[DML EP] Contrib Op: FusedMatMul (#12898)

* Contrib Op: FusedMatMul for DML EP

* Added relevant comments and extra validation

* Polish

* More polish

* Last polish

* Addressed comment on the PR

* Addressed comment on the R

* Removed un-necessary comments

* Used c++ standard function

* used std::c++ algorithms function

* Removed unsed code

Co-authored-by: Sumit Agarwal <sumitagarwal@microsoft.com>
Co-authored-by: Dwayne Robinson <fdwr@hotmail.com>
This commit is contained in:
sumitsays 2022-09-09 09:37:38 -07:00 committed by GitHub
parent eb6aa861cf
commit 05c65a54b3
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
20 changed files with 561 additions and 115 deletions

View file

@ -25,7 +25,7 @@ using ONNX_NAMESPACE::OPTIONAL_VALUE;
void RegisterDmlSchemas() {
MS_DML_OPERATOR_SCHEMA(FusedConv)
MS_DML_OPERATOR_SCHEMA(DmlFusedConv)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused Conv+Activation)DOC")
@ -52,7 +52,7 @@ void RegisterDmlSchemas() {
ONNX_NAMESPACE::convPoolShapeInference(ctx, true, false, 0, 1);
});
MS_DML_OPERATOR_SCHEMA(FusedConvTranspose)
MS_DML_OPERATOR_SCHEMA(DmlFusedConvTranspose)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused ConvTranspose+Activation)DOC")
@ -79,7 +79,7 @@ void RegisterDmlSchemas() {
.TypeAndShapeInferenceFunction(
[](ONNX_NAMESPACE::InferenceContext& ctx) { ONNX_NAMESPACE::convTransposeShapeInference(ctx); });
MS_DML_OPERATOR_SCHEMA(FusedInstanceNormalization)
MS_DML_OPERATOR_SCHEMA(DmlFusedInstanceNormalization)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused InstanceNormalization+Activation)DOC")
@ -100,7 +100,7 @@ void RegisterDmlSchemas() {
ONNX_NAMESPACE::propagateShapeAndTypeFromFirstInput(ctx);
});
MS_DML_OPERATOR_SCHEMA(FusedBatchNormalization)
MS_DML_OPERATOR_SCHEMA(DmlFusedBatchNormalization)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused BatchNormalization+Activation)DOC")
@ -133,7 +133,7 @@ void RegisterDmlSchemas() {
// the other outputs as well.
});
MS_DML_OPERATOR_SCHEMA(FusedMeanVarianceNormalization)
MS_DML_OPERATOR_SCHEMA(DmlFusedMeanVarianceNormalization)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused MeanVarianceNormalization+Activation)DOC")
@ -151,7 +151,7 @@ void RegisterDmlSchemas() {
.Attr(AttrName::FusedRatio, "", onnx::AttributeProto::FLOAT, false)
.TypeAndShapeInferenceFunction(ONNX_NAMESPACE::propagateShapeAndTypeFromFirstInput);
MS_DML_OPERATOR_SCHEMA(FusedGemm)
MS_DML_OPERATOR_SCHEMA(DmlFusedGemm)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused Gemm+Activation)DOC")
@ -194,7 +194,7 @@ void RegisterDmlSchemas() {
}
});
MS_DML_OPERATOR_SCHEMA(FusedMatMul)
MS_DML_OPERATOR_SCHEMA(DmlFusedMatMul)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused MatMul+Activation)DOC")
@ -283,7 +283,7 @@ void RegisterDmlSchemas() {
resultShape;
});
MS_DML_OPERATOR_SCHEMA(FusedAdd)
MS_DML_OPERATOR_SCHEMA(DmlFusedAdd)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused Add+Activation)DOC")
@ -307,7 +307,7 @@ void RegisterDmlSchemas() {
*ctx.getOutputType(0)->mutable_tensor_type()->mutable_shape());
});
MS_DML_OPERATOR_SCHEMA(FusedSum)
MS_DML_OPERATOR_SCHEMA(DmlFusedSum)
.SetDomain(kMSDmlDomain)
.SinceVersion(1)
.SetDoc(R"DOC(DirectML fused Sum+Activation)DOC")

View file

@ -64,7 +64,7 @@ namespace Dml
auto kernelInputIndices = ReplaceUnusedEdgeIndicesWithSentinel(m_kernelInputIndices);
properties.dmlInputCount = static_cast<uint32_t>(kernelInputIndices.size());
properties.kernelInputIndices = kernelInputIndices.data();
auto kernelOutputIndices = ReplaceUnusedEdgeIndicesWithSentinel(m_kernelOutputIndices);
properties.dmlOutputCount = static_cast<uint32_t>(kernelOutputIndices.size());
properties.kernelOutputIndices = kernelOutputIndices.data();
@ -88,7 +88,7 @@ namespace Dml
m_persistentResourceBinding = DML_BUFFER_BINDING{ m_persistentResource.Get(), 0, persistentResourceSize };
}
std::vector<DML_BUFFER_BINDING> initializationInputBindings(m_kernelInputIndices.size());
ORT_THROW_IF_FAILED(m_executionProvider->InitializeOperator(
@ -183,7 +183,7 @@ namespace Dml
else
{
m_inputTensorDescs.push_back(CreateTensorDescFromInput(
kernelInfo,
kernelInfo,
*m_kernelInputIndices[i],
TensorAxis::DoNotCoerce,
TensorAxis::W,
@ -205,7 +205,7 @@ namespace Dml
else
{
m_outputTensorDescs.push_back(CreateTensorDescFromOutput(
kernelInfo,
kernelInfo,
*m_kernelOutputIndices[i],
TensorAxis::DoNotCoerce,
TensorAxis::W,
@ -216,6 +216,112 @@ namespace Dml
}
}
void DmlOperator::InitializeWithShapes(
const MLOperatorKernelCreationContext& kernelInfo,
const std::optional<const std::vector<std::optional<uint32_t>>>& kernelInputIndices,
const std::optional<const std::vector<std::optional<uint32_t>>>& kernelOutputIndices,
const std::optional<gsl::span<gsl::span<const uint32_t>>> inputShapes,
const std::optional<gsl::span<gsl::span<const uint32_t>>> outputShapes,
uint32_t minDimensionCount
)
{
if (kernelInputIndices)
{
m_kernelInputIndices = *kernelInputIndices;
}
else
{
m_kernelInputIndices.resize(kernelInfo.GetInputCount());
std::iota(m_kernelInputIndices.begin(), m_kernelInputIndices.end(), 0);
}
if (kernelOutputIndices)
{
m_kernelOutputIndices = *kernelOutputIndices;
}
else
{
m_kernelOutputIndices.resize(kernelInfo.GetOutputCount());
std::iota(m_kernelOutputIndices.begin(), m_kernelOutputIndices.end(), 0);
}
for (uint32_t i = 0; i < m_kernelInputIndices.size(); i++)
{
// Update m_kernelInputIndices to reflect optional tensors.
if (m_kernelInputIndices[i] == std::nullopt ||
!kernelInfo.IsInputValid(*m_kernelInputIndices[i]))
{
m_kernelInputIndices[i] = std::nullopt;
m_inputTensorDescs.push_back(TensorDesc());
}
else
{
auto edgeDesc = kernelInfo.GetInputEdgeDescription(*m_kernelInputIndices[i]);
assert(edgeDesc.edgeType == MLOperatorEdgeType::Tensor);
// prioritize the given input shapes
TensorDesc tensorDesc;
if (inputShapes.has_value() && i < (*inputShapes).size())
{
tensorDesc = TensorDesc(
edgeDesc.tensorDataType,
(*inputShapes)[i], // desired
(*inputShapes)[i], // original
TensorAxis::DoNotCoerce,
TensorAxis::W,
TensorAxis::RightAligned,
minDimensionCount,
0
);
}
else if (kernelInfo.HasTensorShapeDescription())
{
std::vector<uint32_t> actualTensorShape = kernelInfo.GetTensorShapeDescription().GetInputTensorShape(*m_kernelInputIndices[i]);
tensorDesc = TensorDesc(
edgeDesc.tensorDataType,
actualTensorShape, // desired
actualTensorShape, // original
TensorAxis::DoNotCoerce,
TensorAxis::W,
TensorAxis::RightAligned,
minDimensionCount,
0
);
}
m_inputTensorDescs.push_back(tensorDesc);
}
}
for (uint32_t i = 0; i < m_kernelOutputIndices.size(); i++)
{
// Update m_kernelOutputIndices to reflect optional tensors.
if (m_kernelOutputIndices[i] == std::nullopt ||
!kernelInfo.IsOutputValid(*m_kernelOutputIndices[i]))
{
m_kernelOutputIndices[i] = std::nullopt;
m_outputTensorDescs.push_back(TensorDesc());
}
else
{
std::optional<gsl::span<const uint32_t>> outputShape;
if (outputShapes.has_value() && i < (*outputShapes).size())
{
outputShape = (*outputShapes)[i];
}
m_outputTensorDescs.push_back(CreateTensorDescFromOutput(
kernelInfo,
*m_kernelOutputIndices[i],
TensorAxis::DoNotCoerce,
TensorAxis::W,
TensorAxis::RightAligned,
outputShape,
minDimensionCount
));
}
}
}
void DmlOperator::Compute(const MLOperatorKernelContext& kernelContext)
{
std::vector<IMLOperatorTensor*> inputTensors = GetInputTensorsForExecute(kernelContext);
@ -231,7 +337,7 @@ namespace Dml
bool DmlOperator::AllowHalfPrecisionComputation() const
{
// Most of our operators work with float data, but some do not. In those cases
// no input params are float tensors. This function returns true if the operator
// no input params are float tensors. This function returns true if the operator
// works with at least one float16 tensor and has no tensors of float32 type
bool usesFloat16Tensors = false;
@ -464,7 +570,7 @@ namespace Dml
}
auto outputShape = outputShapeDescription.GetOutputTensorShape(index);
return TensorDesc(
edgeDesc.tensorDataType,
tensorShape ? *tensorShape : outputShape,

View file

@ -17,7 +17,7 @@ namespace Dml
virtual void Compute(const MLOperatorKernelContext& kernelContext);
protected:
ComPtr<IExecutionProvider> m_executionProvider;
ComPtr<IExecutionProvider> m_executionProvider;
ComPtr<IDMLDevice> m_dmlDevice;
// Tensor descs ordered based on index arrays passed to Initialize
@ -43,11 +43,23 @@ namespace Dml
uint32_t minDimensionCount = NchwDimensionCount
);
// This first tries to create TensorDesc with the given input and output shapes, no broadcasting.
// If the shapes are not present, then it will try to create TensorDesc with the shapes from the actual input tensors and shape inference.
// The inputShapes and kernelInputIndices should have same length. Same for outputShapes and kernelOutputIndices.
void InitializeWithShapes(
const MLOperatorKernelCreationContext& kernelInfo,
const std::optional<const std::vector<std::optional<uint32_t>>>& kernelInputIndices = std::nullopt,
const std::optional<const std::vector<std::optional<uint32_t>>>& kernelOutputIndices = std::nullopt,
const std::optional<gsl::span<gsl::span<const uint32_t>>> inputShapes = std::nullopt,
const std::optional<gsl::span<gsl::span<const uint32_t>>> outputShapes = std::nullopt,
uint32_t minDimensionCount = NchwDimensionCount
);
bool AllowHalfPrecisionComputation() const;
DML_EXECUTION_FLAGS GetExecutionFlags() const;
void SetDmlOperatorDesc(
const DML_OPERATOR_DESC& operatorDesc,
const DML_OPERATOR_DESC& operatorDesc,
const MLOperatorKernelCreationContext& kernelInfo
);
@ -55,11 +67,11 @@ namespace Dml
const DML_OPERATOR_DESC& operatorDesc,
const MLOperatorKernelContext& kernelInfo
);
// Tensors ordered based on index arrays passed to Initialize
std::vector<IMLOperatorTensor*> GetInputTensors(const MLOperatorKernelContext& kernelContext);
std::vector<IMLOperatorTensor*> GetOutputTensors(const MLOperatorKernelContext& kernelContext);
// Retrieves the input/output tensors to be supplied to DirectML for execution. These differ from
// Get[Input|Output]Tensors in that they account for the binding requirements of DML, instead of
// unconditionally retrieving all input and output tensors.
@ -106,7 +118,7 @@ namespace Dml
) const;
private:
// For each input or output of the DML kernel, the corresponding input or output of the original
// For each input or output of the DML kernel, the corresponding input or output of the original
// kernel. Entries for unused DML inputs are nullopt.
std::vector<std::optional<uint32_t>> m_kernelInputIndices;
std::vector<std::optional<uint32_t>> m_kernelOutputIndices;

View file

@ -202,6 +202,6 @@ void CALLBACK QueryBatchNormalization(IMLOperatorSupportQueryContextPrivate* con
}
DML_OP_DEFINE_CREATION_FUNCTION(BatchNormalization, DmlOperatorBatchNormalization);
DML_OP_DEFINE_CREATION_FUNCTION(FusedBatchNormalization, DmlOperatorBatchNormalization);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedBatchNormalization, DmlOperatorBatchNormalization);
} // namespace Dml

View file

@ -107,8 +107,8 @@ public:
DML_OP_DEFINE_CREATION_FUNCTION(Conv, DmlOperatorConvolutionTemplate<DML_CONVOLUTION_MODE_CROSS_CORRELATION, DML_CONVOLUTION_DIRECTION_FORWARD>);
DML_OP_DEFINE_CREATION_FUNCTION(ConvTranspose, DmlOperatorConvolutionTemplate<DML_CONVOLUTION_MODE_CROSS_CORRELATION, DML_CONVOLUTION_DIRECTION_BACKWARD>);
DML_OP_DEFINE_CREATION_FUNCTION(FusedConv, DmlOperatorConvolutionTemplate<DML_CONVOLUTION_MODE_CROSS_CORRELATION, DML_CONVOLUTION_DIRECTION_FORWARD>);
DML_OP_DEFINE_CREATION_FUNCTION(FusedConvTranspose, DmlOperatorConvolutionTemplate<DML_CONVOLUTION_MODE_CROSS_CORRELATION, DML_CONVOLUTION_DIRECTION_BACKWARD>);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedConv, DmlOperatorConvolutionTemplate<DML_CONVOLUTION_MODE_CROSS_CORRELATION, DML_CONVOLUTION_DIRECTION_FORWARD>);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedConvTranspose, DmlOperatorConvolutionTemplate<DML_CONVOLUTION_MODE_CROSS_CORRELATION, DML_CONVOLUTION_DIRECTION_BACKWARD>);
DML_OP_DEFINE_CREATION_FUNCTION(ConvTransposeWithDynamicPads, DmlOperatorConvolutionTemplate<DML_CONVOLUTION_MODE_CROSS_CORRELATION, DML_CONVOLUTION_DIRECTION_BACKWARD, true>);
} // namespace Dml

View file

@ -756,7 +756,7 @@ DML_OP_DEFINE_CREATION_FUNCTION(IsInf, DmlOperatorElementwiseIsInf);
DML_OP_DEFINE_CREATION_FUNCTION(Round, DmlOperatorElementwiseRound);
// Fused operators:
DML_OP_DEFINE_CREATION_FUNCTION(FusedAdd, DmlOperatorElementwiseBinary<DML_ELEMENT_WISE_ADD1_OPERATOR_DESC>);
DML_OP_DEFINE_CREATION_FUNCTION(FusedSum, DmlOperatorElementwiseBinaryLoop<DML_ELEMENT_WISE_ADD1_OPERATOR_DESC>);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedAdd, DmlOperatorElementwiseBinary<DML_ELEMENT_WISE_ADD1_OPERATOR_DESC>);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedSum, DmlOperatorElementwiseBinaryLoop<DML_ELEMENT_WISE_ADD1_OPERATOR_DESC>);
} // namespace Dml

View file

@ -0,0 +1,83 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "precomp.h"
namespace Dml
{
class DmlOperatorFusedMatMul : public DmlOperator
{
public:
DmlOperatorFusedMatMul(const MLOperatorKernelCreationContext& kernelInfo)
: DmlOperator(kernelInfo)
{
// FusedMatMul has two inputs, but DML GEMM requires 3 input bindings (a null binding for the C Tensor).
ML_CHECK_VALID_ARGUMENT(kernelInfo.GetInputCount() == 2);
// Need these shapes to apply transpose and
// numpy MatMul's behavior https://docs.scipy.org/doc/numpy-1.13.0/reference/generated/numpy.matmul.html
std::vector<DimensionType> inputShape0 = kernelInfo.GetTensorShapeDescription().GetInputTensorShape(0);
std::vector<DimensionType> inputShape1 = kernelInfo.GetTensorShapeDescription().GetInputTensorShape(1);
std::vector<DimensionType> outputShape = kernelInfo.GetTensorShapeDescription().GetOutputTensorShape(0);
const int32_t transBatchA = kernelInfo.GetOptionalAttribute<int32_t>(AttrName::TransBatchA, 0);
const int32_t transA = kernelInfo.GetOptionalAttribute<int32_t>(AttrName::TransA, 0);
const int32_t transBatchB = kernelInfo.GetOptionalAttribute<int32_t>(AttrName::TransBatchB, 0);
const int32_t transB = kernelInfo.GetOptionalAttribute<int32_t>(AttrName::TransB, 0);
// As of now, CPU FusedMatMul has this extra validation
// https://github.com/microsoft/onnxruntime/blob/main/onnxruntime/core/providers/cpu/math/matmul_helper.h#L72
// Although, DML kernel can work without this validation, but adding this just to be in sync.
if (transBatchA || transBatchB)
{
ML_CHECK_VALID_ARGUMENT(inputShape0.size() > 2 && inputShape0.size() == inputShape1.size(),
"Two inputs should have same rank and rank >= 3 if transBatchA or transBatchB is true");
}
auto [sizesA, stridesA] = OperatorHelper::GetFusedMatMulSizesAndStrides(inputShape0, transBatchA, transA);
auto [sizesB, stridesB] = OperatorHelper::GetFusedMatMulSizesAndStrides(inputShape1, transBatchB, transB);
OperatorHelper::FusedMatMulShapeMapping(sizesA, stridesA, sizesB, stridesB, outputShape);
// At this point, we have manipulated input/output shapes and strides and
// we do not care about actual input shapes present in the model (.onnx file).
// Create the TensorDesc with the manipulated input shapes becuase we don't want incorrect
// broadcasting to be happen inside TensorDesc constructor.
std::vector<std::optional<uint32_t>> inputIndices = { 0, 1, std::nullopt };
gsl::span<const uint32_t> inputShapes[2] = {sizesA, sizesB};
gsl::span<const uint32_t> outputShapes[1] = {outputShape};
DmlOperator::InitializeWithShapes(kernelInfo, inputIndices, std::nullopt, inputShapes, outputShapes, 1);
m_inputTensorDescs[0].SetStrides(stridesA);
m_inputTensorDescs[1].SetStrides(stridesB);
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
std::optional<ActivationOperatorDesc> fusedActivation = FusionHelpers::TryGetFusedActivationDesc(kernelInfo);
DML_OPERATOR_DESC fusedActivationDmlDesc = fusedActivation ? fusedActivation->GetDmlDesc() : DML_OPERATOR_DESC();
const float alpha = kernelInfo.GetOptionalAttribute<float>(AttrName::Alpha, 1.0f);
DML_GEMM_OPERATOR_DESC gemmDesc = {};
gemmDesc.ATensor = &inputDescs[0];
gemmDesc.BTensor = &inputDescs[1];
gemmDesc.CTensor = nullptr;
gemmDesc.OutputTensor = &outputDescs[0];
gemmDesc.TransA = DML_MATRIX_TRANSFORM_NONE;
gemmDesc.TransB = DML_MATRIX_TRANSFORM_NONE;
gemmDesc.Alpha = alpha;
gemmDesc.Beta = 0.0f;
gemmDesc.FusedActivation = fusedActivation ? &fusedActivationDmlDesc : nullptr;
DML_OPERATOR_DESC opDesc = { DML_OPERATOR_GEMM, &gemmDesc };
SetDmlOperatorDesc(opDesc, kernelInfo);
}
};
DML_OP_DEFINE_CREATION_FUNCTION(FusedMatMul, DmlOperatorFusedMatMul);
} // namespace Dml

View file

@ -56,6 +56,6 @@ public:
};
DML_OP_DEFINE_CREATION_FUNCTION(Gemm, DmlOperatorGemm);
DML_OP_DEFINE_CREATION_FUNCTION(FusedGemm, DmlOperatorGemm);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedGemm, DmlOperatorGemm);
} // namespace Dml

View file

@ -63,6 +63,6 @@ public:
};
DML_OP_DEFINE_CREATION_FUNCTION(InstanceNormalization, DmlOperatorInstanceNormalization);
DML_OP_DEFINE_CREATION_FUNCTION(FusedInstanceNormalization, DmlOperatorInstanceNormalization);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedInstanceNormalization, DmlOperatorInstanceNormalization);
} // namespace Dml

View file

@ -56,6 +56,6 @@ public:
DML_OP_DEFINE_CREATION_FUNCTION(MatMul, DmlOperatorMatMul);
DML_OP_DEFINE_CREATION_FUNCTION(FusedMatMul, DmlOperatorMatMul);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedMatMul, DmlOperatorMatMul);
} // namespace Dml

View file

@ -55,6 +55,6 @@ public:
};
DML_OP_DEFINE_CREATION_FUNCTION(MeanVarianceNormalization, DmlOperatorMeanVarNormalization);
DML_OP_DEFINE_CREATION_FUNCTION(FusedMeanVarianceNormalization, DmlOperatorMeanVarNormalization);
DML_OP_DEFINE_CREATION_FUNCTION(DmlFusedMeanVarianceNormalization, DmlOperatorMeanVarNormalization);
} // namespace Dml

View file

@ -191,6 +191,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(ParametricSoftplus);
DML_OP_EXTERN_CREATION_FUNCTION(Affine);
DML_OP_EXTERN_CREATION_FUNCTION(Dropout);
DML_OP_EXTERN_CREATION_FUNCTION(MatMul);
DML_OP_EXTERN_CREATION_FUNCTION(FusedMatMul);
DML_OP_EXTERN_CREATION_FUNCTION(Cast);
DML_OP_EXTERN_CREATION_FUNCTION(CastLike15);
DML_OP_EXTERN_CREATION_FUNCTION(MemcpyFromHost);
@ -205,15 +206,15 @@ DML_OP_EXTERN_CREATION_FUNCTION(Tan);
DML_OP_EXTERN_CREATION_FUNCTION(Acos);
DML_OP_EXTERN_CREATION_FUNCTION(Asin);
DML_OP_EXTERN_CREATION_FUNCTION(Atan);
DML_OP_EXTERN_CREATION_FUNCTION(FusedConv);
DML_OP_EXTERN_CREATION_FUNCTION(FusedConvTranspose);
DML_OP_EXTERN_CREATION_FUNCTION(FusedInstanceNormalization);
DML_OP_EXTERN_CREATION_FUNCTION(FusedBatchNormalization);
DML_OP_EXTERN_CREATION_FUNCTION(FusedMeanVarianceNormalization);
DML_OP_EXTERN_CREATION_FUNCTION(FusedGemm);
DML_OP_EXTERN_CREATION_FUNCTION(FusedMatMul);
DML_OP_EXTERN_CREATION_FUNCTION(FusedAdd);
DML_OP_EXTERN_CREATION_FUNCTION(FusedSum);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedConv);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedConvTranspose);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedInstanceNormalization);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedBatchNormalization);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedMeanVarianceNormalization);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedGemm);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedMatMul);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedAdd);
DML_OP_EXTERN_CREATION_FUNCTION(DmlFusedSum);
DML_OP_EXTERN_CREATION_FUNCTION(QuantizeLinear);
DML_OP_EXTERN_CREATION_FUNCTION(DequantizeLinear);
DML_OP_EXTERN_CREATION_FUNCTION(Sign);
@ -333,12 +334,12 @@ constexpr static std::array<SupportedTensorDataTypes, 4> supportedTypeListQLinea
SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8,
SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8,
SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8,
SupportedTensorDataTypes::Int32
SupportedTensorDataTypes::Int32
};
constexpr static std::array<SupportedTensorDataTypes, 2> supportedTypeListDynamicQuantizeLinear = {
SupportedTensorDataTypes::Float32,
SupportedTensorDataTypes::Float32,
SupportedTensorDataTypes::UInt8,
};
@ -351,28 +352,28 @@ constexpr auto requiredConstantCpuInputs(Args... args)
// Define a single row of OperatorRegistrationInformation.
#define REG_INFO(version, operatorName, ...) \
#operatorName, OnnxOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kOnnxDomain, Create##operatorName, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName>, false, ##__VA_ARGS__,
#operatorName, OnnxOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kOnnxDomain, Create##operatorName, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName>, false, ##__VA_ARGS__,
// Versioned operator
#define REG_INFO_VER(version, operatorName, ...) \
#operatorName, OnnxOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kOnnxDomain, Create##operatorName##version, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName##version>, false, ##__VA_ARGS__,
#operatorName, OnnxOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kOnnxDomain, Create##operatorName##version, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName##version>, false, ##__VA_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_COPY(version, operatorName, ...) \
#operatorName, OnnxOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kOnnxDomain, CreateCopy, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName##version>, true, ##__VA_ARGS__,
#operatorName, OnnxOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kOnnxDomain, CreateCopy, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName##version>, true, ##__VA_ARGS__,
// MS-domain operators
#define REG_INFO_MS(version, operatorName, ...) \
#operatorName, MsftOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kMSDomain, Create##operatorName, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName>, false, ##__VA_ARGS__,
#operatorName, MsftOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kMSDomain, Create##operatorName, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName>, false, ##__VA_ARGS__,
// MS-domain operators
#define REG_INFO_MSDML(version, operatorName, ...) \
#operatorName, MsftOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kMSDmlDomain, Create##operatorName, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName>, false, ##__VA_ARGS__,
#operatorName, MsftOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kMSDmlDomain, Create##operatorName, ShapeInferenceFunction<ShapeInferenceHelper_##operatorName>, false, ##__VA_ARGS__,
constexpr static OperatorRegistrationInformation operatorRegistrationInformationTable[] =
{
/// Domain/Type, Ver, Name, TypeNames, Types, Graph Support, Required const CPU inputs,
/// Domain/Type, Ver, Name, TypeNames, Types, Graph Support, Required const CPU inputs,
/// Input count required for graph support,
/// Support query function
@ -657,7 +658,6 @@ 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)},
@ -677,17 +677,21 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation
// 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)},
{REG_INFO_MSDML(1, FusedConvTranspose, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, FusedInstanceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, FusedBatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, FusedMeanVarianceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, FusedGemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, FusedMatMul, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, FusedAdd, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, FusedSum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)},
// DmlFused operators
{REG_INFO_MSDML(1, DmlFusedConv, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, DmlFusedConvTranspose, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, DmlFusedInstanceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, DmlFusedBatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, DmlFusedMeanVarianceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, DmlFusedGemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, DmlFusedMatMul, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, DmlFusedAdd, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MSDML(1, DmlFusedSum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)},
// Contrib operators
{REG_INFO_MS( 1, Gelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MS( 1, FusedMatMul, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO( 10, IsInf, typeNameListTwo, supportedTypeListIsInf, DmlGraphSupport::Supported)},
{REG_INFO( 10, Mod, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported)},
{REG_INFO( 13, Mod, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported)},
@ -700,7 +704,7 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation
{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_MS( 1, QLinearAdd, typeNameListDefault, supportedTypeListInteger8, DmlGraphSupport::Supported)},
{REG_INFO( 10, QLinearConv, typeNameListFour, supportedTypeListQLinearConv, DmlGraphSupport::Supported)},
{REG_INFO( 10, QLinearMatMul, typeNameListThree, supportedTypeListQLinearMatMul, DmlGraphSupport::Supported)},
@ -708,8 +712,8 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation
{REG_INFO( 10, ConvInteger, typeNameListThree, supportedTypeListInteger, DmlGraphSupport::Supported)},
{REG_INFO( 11, DynamicQuantizeLinear, typeNameListTwo, supportedTypeListDynamicQuantizeLinear, DmlGraphSupport::Supported)},
};
template<typename T>
template<typename T>
MLOperatorEdgeDescription EdgeDesc()
{
return {MLOperatorEdgeType::Tensor, static_cast<uint64_t>(MLTypeTraits<T>::TensorType)};
@ -726,7 +730,7 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry)
for (const OperatorRegistrationInformation& information : operatorRegistrationInformationTable)
{
assert(information.tensorTypeNames.size() == information.supportedTensorDataTypes.size());
MLOperatorKernelDescription desc = {};
desc.domain = information.domain;
desc.name = information.operatorName;
@ -735,11 +739,11 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry)
// The graph must be configured with operators from only the legacy DML API, or only the new DML API
bool kernelSupportsGraph = !bool(information.dmlGraphSupport & DmlGraphSupport::NotSupported);
desc.options = information.shapeInferenceFunction ?
desc.options = information.shapeInferenceFunction ?
MLOperatorKernelOptions::None : MLOperatorKernelOptions::AllowDynamicInputShapes;
desc.minimumOperatorSetVersion = information.sinceVersion;
typeConstraints.resize(information.tensorTypeNames.size());
desc.typeConstraints = typeConstraints.data();
desc.typeConstraintCount = static_cast<uint32_t>(typeConstraints.size());
@ -750,7 +754,7 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry)
FusionHelpers::AssertFusableOperatorSupportsVersionIfExists(desc.name, desc.domain, desc.minimumOperatorSetVersion);
#endif
// edgeDescs will accumulate the edge descriptions across all type constraints.
// edgeDescs will accumulate the edge descriptions across all type constraints.
// The values of allowedTypeCount will indicate how many elements of edgeDescs
// belong to each type constraint.
edgeDescs.clear();
@ -773,7 +777,7 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry)
if (bool(supportedTypes & SupportedTensorDataTypes::Int64 )) edgeDescs.push_back(EdgeDesc<int64_t>());
//if (bool(supportedTypes & SupportedTensorDataTypes::String )) edgeDescs.push_back(EdgeDesc<std::string>());
if (bool(supportedTypes & SupportedTensorDataTypes::Bool )) edgeDescs.push_back(EdgeDesc<bool>());
if (bool(supportedTypes & SupportedTensorDataTypes::Float16)) edgeDescs.push_back(EdgeDesc<::MLFloat16>());
if (bool(supportedTypes & SupportedTensorDataTypes::Float16)) edgeDescs.push_back(EdgeDesc<::MLFloat16>());
if (bool(supportedTypes & SupportedTensorDataTypes::Float64)) edgeDescs.push_back(EdgeDesc<double>());
if (bool(supportedTypes & SupportedTensorDataTypes::UInt32 )) edgeDescs.push_back(EdgeDesc<uint32_t>());
if (bool(supportedTypes & SupportedTensorDataTypes::UInt64 )) edgeDescs.push_back(EdgeDesc<uint64_t>());
@ -781,7 +785,7 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry)
typeConstraints[i].allowedTypeCount = static_cast<uint32_t>(edgeDescs.size() - lastEdgeDescSize);
lastEdgeDescSize = edgeDescs.size();
}
// Now that the edge descriptions list won't be re-allocated, assign pointers to its memory
// into the type constraints entries
size_t totalTypeCount = 0;
@ -793,7 +797,7 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry)
ComPtr<MLOperatorKernelFactory> factory = wil::MakeOrThrow<MLOperatorKernelFactory>(information.creationFunction);
ComPtr<MLOperatorShapeInferrer> shapeInferrer;
if (information.shapeInferenceFunction)
{
shapeInferrer = wil::MakeOrThrow<MLOperatorShapeInferrer>(information.shapeInferenceFunction);
@ -806,8 +810,8 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry)
}
ORT_THROW_IF_FAILED(registryPrivate->RegisterOperatorKernel(
&desc,
factory.Get(),
&desc,
factory.Get(),
shapeInferrer.Get(),
supportQuery.Get(),
true, // isInternalOperator

View file

@ -123,7 +123,7 @@ namespace Dml
}
}
} // namespace ActivationHelper
namespace FusionHelpers
{
struct OperatorInfo
@ -225,7 +225,7 @@ namespace Dml
return std::nullopt;
}
if (!opIt->activationFilter.empty() &&
if (!opIt->activationFilter.empty() &&
std::find(opIt->activationFilter.begin(), opIt->activationFilter.end(), activationOpType) == opIt->activationFilter.end())
{
return std::nullopt;
@ -236,8 +236,8 @@ namespace Dml
return std::nullopt;
}
// All fused ops just have "Fused" prepended to their name (e.g. "Conv" -> "FusedConv").
std::string fusedOpType = std::string("Fused").append(candidateOpType);
// All fused ops have "DmlFused" prepended to their name (e.g. "Conv" -> "DmlFusedConv").
std::string fusedOpType = std::string("DmlFused").append(candidateOpType);
return FusedOpProperties{ std::move(fusedOpType), onnxruntime::kMSDmlDomain };
}

View file

@ -32,9 +32,9 @@ TensorDesc::TensorDesc(
m_bufferTensorDesc.GuaranteedBaseOffsetAlignment = guaranteedBaseOffsetAlignment;
m_bufferTensorDesc.TotalTensorSizeInBytes = DMLCalcBufferTensorSize(
m_bufferTensorDesc.DataType,
m_bufferTensorDesc.DimensionCount,
m_sizes,
m_bufferTensorDesc.DataType,
m_bufferTensorDesc.DimensionCount,
m_sizes,
strides ? m_strides : nullptr
);
}
@ -210,6 +210,16 @@ gsl::span<const uint32_t> TensorDesc::GetStrides() const
return { m_strides, m_strides + m_bufferTensorDesc.DimensionCount };
}
void TensorDesc::SetStrides(gsl::span<const uint32_t> strides)
{
if (!strides.empty())
{
ML_CHECK_VALID_ARGUMENT(strides.size() <= std::size(m_strides));
m_bufferTensorDesc.Strides = strides.data();
std::copy(strides.begin(), strides.end(), m_strides);
}
}
DML_TENSOR_DESC TensorDesc::GetDmlDesc()
{
if (m_tensorType == DML_TENSOR_TYPE_INVALID)
@ -217,6 +227,8 @@ DML_TENSOR_DESC TensorDesc::GetDmlDesc()
return { m_tensorType, nullptr };
}
// Update the DML_BUFFER_TENSOR_DESC Sizes and Strides pointers to point internally to the TensorDesc fields.
// This update matters whether it was a new instance or a copy from via copy constructor from another TensorDesc.
m_bufferTensorDesc.Sizes = m_sizes;
if (m_bufferTensorDesc.Strides)
{

View file

@ -43,6 +43,7 @@ namespace Dml
void SetDimensionCount(uint32_t newDimensionCount, TensorAxis alignment);
gsl::span<const uint32_t> GetSizes() const { return { m_sizes, m_sizes + m_bufferTensorDesc.DimensionCount }; }
gsl::span<const uint32_t> GetStrides() const;
void SetStrides(gsl::span<const uint32_t> strides);
inline uint64_t GetBufferSizeInBytes() const
{

View file

@ -90,7 +90,9 @@ namespace AttrName
static constexpr const char* To = "to";
static constexpr const char* TrainingMode = "training_mode";
static constexpr const char* TransA = "transA";
static constexpr const char* TransBatchA = "transBatchA";
static constexpr const char* TransB = "transB";
static constexpr const char* TransBatchB = "transBatchB";
static constexpr const char* Upper = "upper";
static constexpr const char* Value = "value";
static constexpr const char* WidthScale = "width_scale";

View file

@ -332,7 +332,7 @@ namespace OperatorHelper
{
uint32_t padding = args.startPadding[dim] + args.endPadding[dim];
uint32_t kernelLength = 1 + (args.windowSize[dim] - 1) * args.dilations[dim];
outputDimensions[dimOffset + dim] = (inputDimensions[dimOffset + dim] - 1) * args.strides[dim] + kernelLength + args.outputPadding[dim] - padding;
}
@ -460,7 +460,7 @@ namespace OperatorHelper
{
std::vector<int> outputPadding = kernelInfo.GetOptionalAttributeVectorInt32(AttrName::OutputPadding);
ML_CHECK_VALID_ARGUMENT(outputPadding.size() >= 2);
std::copy(outputPadding.begin(), outputPadding.begin() + spatialDimensionCount, args.outputPadding);
}
else
@ -522,7 +522,7 @@ namespace OperatorHelper
args.endPadding[dim] = padding - args.startPadding[dim];
}
}
void MatMulShapeMapping(
std::vector<DimensionType>& inputShape0,
std::vector<DimensionType>& inputShape1,
@ -550,6 +550,81 @@ namespace OperatorHelper
inputShape1.insert(inputShape1.begin(), outputShape.begin(), outputShape.end() - 2);
}
void FusedMatMulShapeMapping(
std::vector<DimensionType>& inputShape0,
std::vector<DimensionType>& inputStride0,
std::vector<DimensionType>& inputShape1,
std::vector<DimensionType>& inputStride1,
std::vector<DimensionType>& outputShape)
{
// Get the padded input shapes, and undo the effect of padding removal from the output shape.
if (inputShape1.size() == 1)
{
inputShape1.push_back(1);
inputStride1.push_back(0);
outputShape.push_back(1);
}
if (inputShape0.size() == 1)
{
inputShape0.insert(inputShape0.begin(), 1);
inputStride0.insert(inputStride0.begin(), 0);
outputShape.insert(outputShape.end() - 1, 1);
}
auto broadcastedRank = std::max(inputShape0.size(), inputShape1.size());
inputShape0.insert(inputShape0.begin(), (broadcastedRank - inputShape0.size()), 1);
inputStride0.insert(inputStride0.begin(), (broadcastedRank - inputStride0.size()), 0);
inputShape1.insert(inputShape1.begin(), (broadcastedRank - inputShape1.size()), 1);
inputStride1.insert(inputStride1.begin(), (broadcastedRank - inputStride1.size()), 0);
BroadcastTensorShapeAndSetStrides(
gsl::make_span(inputShape0.data(), broadcastedRank - 2),
gsl::make_span(inputStride0.data(), broadcastedRank - 2),
gsl::make_span(inputShape1.data(), broadcastedRank - 2),
gsl::make_span(inputStride1.data(), broadcastedRank - 2)
);
}
std::pair<std::vector<uint32_t>, std::vector<uint32_t>> GetFusedMatMulSizesAndStrides(
gsl::span<const uint32_t> sizes,
int32_t transBatch,
int32_t transpose)
{
const uint32_t dimensionCount = sizes.size();
std::vector<uint32_t> newStrides(dimensionCount);
std::vector<uint32_t> newSizes(sizes.begin(), sizes.end());
// Calculate packed strides.
uint32_t stride = 1;
for (int i = dimensionCount - 1; i >= 0; i--)
{
newStrides[i] = stride;
stride *= sizes[i];
}
// According to contrib ops shape inference
// https://github.com/microsoft/onnxruntime/blob/main/onnxruntime/core/graph/contrib_ops/contrib_defs.cc#L215
// `transBatch` needs to be applied first and then `transpose`.
if (transBatch)
{
ML_CHECK_VALID_ARGUMENT(dimensionCount > 2,
"FusedMatMul operator: Tensor size should be more than 2, if attribute transBatch is true");
std::rotate(newSizes.begin(), newSizes.end() - 2, newSizes.end() - 1);
std::rotate(newStrides.begin(), newStrides.end() - 2, newStrides.end() - 1);
}
if (transpose)
{
std::swap(newStrides[dimensionCount - 2], newStrides[dimensionCount - 1]);
std::swap(newSizes[dimensionCount - 2], newSizes[dimensionCount - 1]);
}
return std::make_pair(newSizes, newStrides);
}
std::vector<EdgeShapes> GetOutputShapeAsInputShapeHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
assert(shapeInfo.GetInputCount() > m_inputTensorIndex);
@ -617,18 +692,73 @@ namespace OperatorHelper
return outputShape;
}
void BroadcastTensorShapeAndSetStrides(
gsl::span<DimensionType> inputShape0,
gsl::span<DimensionType> inputStride0,
gsl::span<DimensionType> inputShape1,
gsl::span<DimensionType> inputStride1
)
{
if (inputShape0 != inputShape1)
{
ML_CHECK_VALID_ARGUMENT(
inputShape0.size() == inputShape1.size() &&
inputShape0.size() == inputStride0.size() &&
inputStride0.size() == inputStride1.size(),
"Size of inputShape0, inputStride0, inputShape1 and inputStride1 should be same while broadcasting");
// Walk backwards through both input shapes and broadcast each dimension,
// ignoring the last 2 dimensions (matrix dimensions).
auto rank = inputShape0.size();
auto inDim0Iter = inputShape0.rbegin();
auto inDim1Iter = inputShape1.rbegin();
auto inStride0Iter = inputStride0.rbegin();
auto inStride1Iter = inputStride1.rbegin();
while (rank-- > 0)
{
DimensionType inDimension0 = *inDim0Iter;
DimensionType inStride0 = *inStride0Iter;
DimensionType inDimension1 = *inDim1Iter;
DimensionType inStride1 = *inStride1Iter;
// 0-sized dimensions indicate an empty tensor and shouldn't be broadcasted to higher dimensions.
if (inDimension0 == 0 || inDimension1 == 0)
{
inDimension0 = 0;
inDimension1 = 0;
}
ML_CHECK_VALID_ARGUMENT((inDimension0 == inDimension1) || (inDimension0 == 1) || (inDimension1 == 1));
auto broadcastedDimension = std::max(inDimension0, inDimension1);
inputShape0[rank] = broadcastedDimension;
inputShape1[rank] = broadcastedDimension;
inputStride0[rank] = (broadcastedDimension != inDimension0) ? 0 : inStride0;
inputStride1[rank] = (broadcastedDimension != inDimension1) ? 0 : inStride1;
++inDim0Iter;
++inStride0Iter;
++inDim1Iter;
++inStride1Iter;
}
}
}
void ConvolutionHelperBase::InitializeKernelAndShapes(const IShapeInformationAdapter& shapeInformation)
{
const std::vector<DimensionType> inputDimensions = shapeInformation.GetInputTensorShape(m_inputTensorIndex);
const std::vector<DimensionType> filterDims = shapeInformation.GetInputTensorShape(m_filterTensorIndex);
ML_CHECK_VALID_ARGUMENT(
inputDimensions.size() >= 3 && inputDimensions.size() <= 5,
"Input dimensions must be: 3, 4, 5."
);
ResolvingPadding(inputDimensions);
m_outputShapes.resize(1);
m_outputShapes[0] = InitializeKernelOutputDimensions(inputDimensions, m_kernel);
m_outputShapes[0].GetShape()[C] = filterDims[K];
@ -966,7 +1096,7 @@ namespace OperatorHelper
m_startPadding.size() == outputDimensions.size() &&
m_endPadding.size() == outputDimensions.size()
);
for (size_t i = 0; i < outputDimensions.size(); ++i)
{
outputDimensions[i] += m_startPadding[i] + m_endPadding[i];
@ -1463,8 +1593,8 @@ namespace OperatorHelper
bool EinSumHelper::IsMatMulOperatorType() const noexcept
{
return m_recognizedOperatorType == RecognizedOperatorType::MatMul ||
m_recognizedOperatorType == RecognizedOperatorType::MatMulTransposeA ||
return m_recognizedOperatorType == RecognizedOperatorType::MatMul ||
m_recognizedOperatorType == RecognizedOperatorType::MatMulTransposeA ||
m_recognizedOperatorType == RecognizedOperatorType::MatMulTransposeB;
}
@ -1522,6 +1652,71 @@ namespace OperatorHelper
return {std::move(outputDims)};
}
std::vector<EdgeShapes> FusedMatMulHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
ML_CHECK_VALID_ARGUMENT(shapeInfo.GetInputCount() == 2);
// Following numpy.matmul for shape inference:
// https://docs.scipy.org/doc/numpy/reference/generated/numpy.matmul.html
// The behavior depends on the arguments in the following way.
// * If both arguments are 2 - D they are multiplied like conventional matrices.
// * If either argument is N - D, N > 2, it is treated as a stack of matrices residing in the last two indexes and broadcast accordingly.
// * If the first argument is 1 - D, it is promoted to a matrix by prepending a 1 to its dimensions. After matrix multiplication the prepended 1 is removed.
// * If the second argument is 1 - D, it is promoted to a matrix by appending a 1 to its dimensions. After matrix multiplication the appended 1 is removed.
auto inputShape0 = shapeInfo.GetInputTensorShape(0);
auto inputShape1 = shapeInfo.GetInputTensorShape(1);
ML_CHECK_VALID_ARGUMENT(inputShape0.size() >= 1);
ML_CHECK_VALID_ARGUMENT(inputShape1.size() >= 1);
auto [sizesA, stridesA] = GetFusedMatMulSizesAndStrides(
inputShape0,
shapeInfo.GetOptionalAttribute(AttrName::TransBatchA, -1),
shapeInfo.GetOptionalAttribute(AttrName::TransA, -1)
);
inputShape0 = sizesA;
auto [sizesB, stridesB] = GetFusedMatMulSizesAndStrides(
inputShape1,
shapeInfo.GetOptionalAttribute(AttrName::TransBatchB, -1),
shapeInfo.GetOptionalAttribute(AttrName::TransB, -1)
);
inputShape1 = sizesB;
std::vector<uint32_t> outputMatrixDims;
// Modify the input and truncated output shapes per the above comments.
// The extra dimensions of the output beyond the two matrix dimensions
// will be computed afterward by broadcasting.
if (inputShape0.size() == 1)
{
inputShape0.insert(inputShape0.begin(), 1);
}
else
{
outputMatrixDims.push_back(inputShape0[inputShape0.size() - 2]);
}
if (inputShape1.size() == 1)
{
inputShape1.push_back(1);
}
else
{
outputMatrixDims.push_back(inputShape1[inputShape1.size() - 1]);
}
// Remove the matrix dimensions from each input, resulting in broadcastable shapes.
std::vector<uint32_t> batchDims0(inputShape0.begin(), inputShape0.end() - 2);
std::vector<uint32_t> batchDims1(inputShape1.begin(), inputShape1.end() - 2);
// Broadcast the extra dimensions of each input, then add the truncated matrix dimensions.
std::vector<uint32_t> outputDims = BroadcastTensorShape(batchDims0, batchDims1);
outputDims.insert(outputDims.end(), outputMatrixDims.begin(), outputMatrixDims.end());
return {std::move(outputDims)};
}
void TopKHelper::Initialize(
const IKernelInformationAdapter& kernelInformation,
const IShapeInformationAdapter& shapeInformation,
@ -1611,7 +1806,7 @@ namespace OperatorHelper
}
}
}
std::vector<EdgeShapes> RandomUniformHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
EdgeShapes outputShape(m_tensorShape);
@ -1737,7 +1932,7 @@ namespace OperatorHelper
// the first dimension of the output. The default axis is 1, which
// preserves the first dimension (typically batch size), and flattens
// the remaining dimensions. An axis of 0 means no input dimensions
// are flattened into the first output dimension, so the output is 1D
// are flattened into the first output dimension, so the output is 1D
// padded with a 1 in the first diemension.
ML_CHECK_VALID_ARGUMENT(m_axis >= 0 && m_axis <= gsl::narrow_cast<int>(inputDimensions.size()));
gsl::span<const DimensionType> outputDimensionsSpan(inputDimensions);
@ -1770,7 +1965,7 @@ namespace OperatorHelper
auto roiShape = shapeInfo.GetInputTensorShape(InputTensors::ROIS);
auto inputShape = shapeInfo.GetInputTensorShape(InputTensors::INPUT);
ML_CHECK_VALID_ARGUMENT(inputShape.size() >= 4, "inputShape must be >= 4.");
DimensionType outputDimensions[4] =
{
roiShape[0], // number of ROIs
@ -1926,7 +2121,7 @@ namespace OperatorHelper
return { outputDimensions };
}
std::vector<EdgeShapes> SpaceToDepthHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
auto inputDimensions = shapeInfo.GetInputTensorShape(0);
@ -2010,7 +2205,7 @@ namespace OperatorHelper
outputDimensions[inferDim] = inElementCount / outElementCount;
}
}
return { EdgeShapes(outputDimensions) };
}
@ -2036,10 +2231,10 @@ namespace OperatorHelper
// Determine the broadcasted input shape.
outputDimensions = OperatorHelper::BroadcastTensorShape(actualInputTensorShape, desiredTensorShape);
return { EdgeShapes(outputDimensions) };
}
std::vector<EdgeShapes> ConstantOfShapeHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
std::vector<uint32_t> outputDimensions;
@ -2108,7 +2303,7 @@ namespace OperatorHelper
auto& attributes = kernelInformation.GetAttributes();
m_inputDimensions = shapeInformation.GetInputTensorShape(0);
std::vector<int32_t> outputSizes;
if (opsetVersion >= 11)
{
if (kernelInformation.IsInputValid(1))
@ -2139,7 +2334,7 @@ namespace OperatorHelper
// From attribute, compatible with Upsample-7.
m_scales = attributes.template GetOptionalAttribute<std::vector<float>>(AttrName::Scales, std::vector<float>());
}
assert(m_outputDimensions.empty());
ML_CHECK_VALID_ARGUMENT(m_scales.empty() || outputSizes.empty(), "scales and roi cannot both be present.");

View file

@ -18,6 +18,14 @@ std::vector<DimensionType> BroadcastTensorShape(
gsl::span<const DimensionType> inputShape0,
gsl::span<const DimensionType> inputShape1);
// This won't allocate extra memory, if required. This expects
// caller to make the size of all containers to be same
void BroadcastTensorShapeAndSetStrides(
gsl::span<DimensionType> inputShape0,
gsl::span<DimensionType> inputStride0,
gsl::span<DimensionType> inputShape1,
gsl::span<DimensionType> inputStride1);
// Find all the occurrences of a value, and return the array indices (in ascending order).
//
// e.g. input values = {2,1,3,1,1,5}
@ -218,6 +226,18 @@ void MatMulShapeMapping(
std::vector<DimensionType>& inputShape1,
std::vector<DimensionType>& outputShape);
void FusedMatMulShapeMapping(
std::vector<DimensionType>& inputShape0,
std::vector<DimensionType>& inputStride0,
std::vector<DimensionType>& inputShape1,
std::vector<DimensionType>& inputStride1,
std::vector<DimensionType>& outputShape);
std::pair<std::vector<uint32_t>, std::vector<uint32_t>> GetFusedMatMulSizesAndStrides(
gsl::span<const uint32_t> sizes,
int32_t transBatch = 0,
int32_t transpose = 0);
class GetOutputShapeAsInputShapeHelper
{
public:
@ -762,6 +782,15 @@ public:
MatMulHelper(const Info_t& info, const Shape_t& shape) : MatMulHelperBase(info, shape, 0, 1) {}
};
class FusedMatMulHelper
{
public:
template<typename Info_t, typename Shape_t>
FusedMatMulHelper(const Info_t& info, const Shape_t& shape) {}
std::vector<EdgeShapes> GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const;
};
class QLinearMatMulHelper : public MatMulHelperBase
{
public:
@ -1519,14 +1548,15 @@ using ShapeInferenceHelper_Range = RangeHelper;
using ShapeInferenceHelper_CastLike15 = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_FusedConv = ConvHelper;
using ShapeInferenceHelper_FusedConvTranspose = ConvTransposeHelper;
using ShapeInferenceHelper_FusedInstanceNormalization = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_FusedBatchNormalization = BatchNormalizationHelper;
using ShapeInferenceHelper_FusedMeanVarianceNormalization = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_FusedGemm = GemmHelper;
using ShapeInferenceHelper_FusedMatMul = MatMulHelper;
using ShapeInferenceHelper_FusedAdd = GetBroadcastedOutputShapeHelper;
using ShapeInferenceHelper_FusedSum = GetBroadcastedOutputShapeHelper;
using ShapeInferenceHelper_DmlFusedConv = ConvHelper;
using ShapeInferenceHelper_DmlFusedConvTranspose = ConvTransposeHelper;
using ShapeInferenceHelper_DmlFusedInstanceNormalization = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_DmlFusedBatchNormalization = BatchNormalizationHelper;
using ShapeInferenceHelper_DmlFusedMeanVarianceNormalization = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_DmlFusedGemm = GemmHelper;
using ShapeInferenceHelper_DmlFusedMatMul = MatMulHelper;
using ShapeInferenceHelper_FusedMatMul = FusedMatMulHelper;
using ShapeInferenceHelper_DmlFusedAdd = GetBroadcastedOutputShapeHelper;
using ShapeInferenceHelper_DmlFusedSum = GetBroadcastedOutputShapeHelper;
} // namespace OperatorHelper

View file

@ -367,20 +367,21 @@ namespace OperatorHelper
namespace MsftOperatorSet1
{
static const int sc_sinceVer_FusedConv = 1;
static const int sc_sinceVer_FusedConvTranspose = 1;
static const int sc_sinceVer_FusedInstanceNormalization = 1;
static const int sc_sinceVer_FusedBatchNormalization = 1;
static const int sc_sinceVer_FusedMeanVarianceNormalization = 1;
static const int sc_sinceVer_FusedGemm = 1;
static const int sc_sinceVer_FusedMatMul = 1;
static const int sc_sinceVer_FusedAdd = 1;
static const int sc_sinceVer_FusedSum = 1;
static const int sc_sinceVer_DmlFusedConv = 1;
static const int sc_sinceVer_DmlFusedConvTranspose = 1;
static const int sc_sinceVer_DmlFusedInstanceNormalization = 1;
static const int sc_sinceVer_DmlFusedBatchNormalization = 1;
static const int sc_sinceVer_DmlFusedMeanVarianceNormalization = 1;
static const int sc_sinceVer_DmlFusedGemm = 1;
static const int sc_sinceVer_DmlFusedMatMul = 1;
static const int sc_sinceVer_DmlFusedAdd = 1;
static const int sc_sinceVer_DmlFusedSum = 1;
static const int sc_sinceVer_QuantizeLinear = 1;
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;
static const int sc_sinceVer_FusedMatMul = 1;
} // namespace MsftOperatorSet1
} // namespace OperatorHelper

View file

@ -104,11 +104,11 @@ static void CustomOperatorFusion() {
REGISTER_KERNEL(Conv, onnxruntime::kOnnxDomain, OnnxOperatorSet7, ConvHelper, &m_callCounts.conv);
REGISTER_KERNEL(Relu, onnxruntime::kOnnxDomain, OnnxOperatorSet7, GetOutputShapeAsInputShapeHelper, &m_callCounts.relu);
REGISTER_KERNEL(FusedConv, onnxruntime::kMSDmlDomain, MsftOperatorSet1, ConvHelper, &m_callCounts.fusedConv);
REGISTER_KERNEL(DmlFusedConv, onnxruntime::kMSDmlDomain, MsftOperatorSet1, ConvHelper, &m_callCounts.fusedConv);
REGISTER_KERNEL(Gemm, onnxruntime::kOnnxDomain, OnnxOperatorSet7, GemmHelper, &m_callCounts.gemm);
REGISTER_KERNEL(Sigmoid, onnxruntime::kOnnxDomain, OnnxOperatorSet7, GetOutputShapeAsInputShapeHelper, &m_callCounts.sigmoid);
REGISTER_KERNEL(FusedGemm, onnxruntime::kMSDmlDomain, MsftOperatorSet1, GemmHelper, &m_callCounts.fusedGemm);
REGISTER_KERNEL(DmlFusedGemm, onnxruntime::kMSDmlDomain, MsftOperatorSet1, GemmHelper, &m_callCounts.fusedGemm);
REGISTER_KERNEL(BatchNormalization, onnxruntime::kOnnxDomain, OnnxOperatorSet7, GetOutputShapeAsInputShapeHelper, &m_callCounts.batchNorm);
REGISTER_KERNEL(MaxPool, onnxruntime::kOnnxDomain, OnnxOperatorSet7, PoolingHelper, &m_callCounts.maxPool);