Merged PR 5100436: EinSum ONNX 1.7 (opset 12) ORT DML EP kernel

Adds EinSum operator (purely an EP kernel, not a dedicated DML operator), which takes an equation string and depending on the specifics is capable of representing: identity, diag, trace, transpose, reduce sum, dot product, matmul, elementwise multiplication, inner product, outer product.

The DML EP recognizes many of them (identity, transpose, reduce sum, 1D dot product, matmul, elementwise multiplication), but defers to CPU when not supported (extended inner product, outer product, diag, trace, arbitrary batch ellipsis).

https://github.com/onnx/onnx/blob/master/docs/Operators.md#Einsum

WindowsAI PR: https://microsoft.visualstudio.com/DefaultCollection/WindowsAI/_git/WindowsAI/pullrequest/5100608

Related work items: #27469790
This commit is contained in:
Dwayne Robinson 2020-08-27 22:10:14 +00:00
parent b5c765c76b
commit f68d5263b7
12 changed files with 531 additions and 40 deletions

View file

@ -132,4 +132,16 @@ uint32_t GetSupportedDeviceDataTypeMask(IDMLDevice* dmlDevice)
return deviceTypeMask;
}
void GetDescendingPackedStrides(gsl::span<const uint32_t> sizes, /*out*/ gsl::span<uint32_t> strides)
{
assert(sizes.size() == strides.size());
uint32_t stride = 1;
for (size_t i = strides.size(); i-- > 0; )
{
strides[i] = stride;
stride *= sizes[i];
}
}
} // namespace Dml

View file

@ -19,6 +19,7 @@ namespace Dml
size_t ComputeByteSizeFromDimensions(gsl::span<const DimensionType> dimensions, MLOperatorTensorDataType tensorDataType);
size_t ComputeByteSizeFromTensor(IMLOperatorTensor& tensor);
uint32_t GetSupportedDeviceDataTypeMask(IDMLDevice* dmlDevice);
void GetDescendingPackedStrides(gsl::span<const uint32_t> sizes, /*out*/ gsl::span<uint32_t> strides);
bool IsSigned(DML_TENSOR_DATA_TYPE dataType);

View file

@ -0,0 +1,176 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "precomp.h"
namespace Dml
{
class DmlOperatorEinSum : public DmlOperator, public EinSumHelper
{
public:
DmlOperatorEinSum(const MLOperatorKernelCreationContext& kernelCreationContext, uint32_t opsetVersion)
: DmlOperator(kernelCreationContext),
EinSumHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription(), opsetVersion)
{
ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() + 1 == m_components.size(), "EinSum input tensor count is inconsistent with the equation component count.");
ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1, "EinSum expects one output tensor.");
DmlOperator::Initialize(kernelCreationContext);
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
static_assert(RecognizedOperatorType::Total == static_cast<RecognizedOperatorType>(8), "Update this switch.");
switch (m_recognizedOperatorType)
{
case RecognizedOperatorType::Multiply:
{
DML_ELEMENT_WISE_MULTIPLY_OPERATOR_DESC operatorDesc = {};
operatorDesc.ATensor = &inputDescs[0];
operatorDesc.BTensor = &inputDescs[1];
operatorDesc.OutputTensor = outputDescs.data();
SetDmlOperatorDesc({ DML_OPERATOR_ELEMENT_WISE_MULTIPLY, &operatorDesc}, kernelCreationContext);
}
break;
case RecognizedOperatorType::MatMul:
case RecognizedOperatorType::MatMulTransposeA:
case RecognizedOperatorType::MatMulTransposeB:
{
DML_GEMM_OPERATOR_DESC operatorDesc = {};
operatorDesc.ATensor = &inputDescs[0];
operatorDesc.BTensor = &inputDescs[1];
// No operatorDesc.CTensor
operatorDesc.OutputTensor = &outputDescs[0];
operatorDesc.TransA = (m_recognizedOperatorType == RecognizedOperatorType::MatMulTransposeA) ? DML_MATRIX_TRANSFORM_TRANSPOSE : DML_MATRIX_TRANSFORM_NONE;
operatorDesc.TransB = (m_recognizedOperatorType == RecognizedOperatorType::MatMulTransposeB) ? DML_MATRIX_TRANSFORM_TRANSPOSE : DML_MATRIX_TRANSFORM_NONE;
operatorDesc.Alpha = 1.0;
operatorDesc.Beta = 0.0;
operatorDesc.FusedActivation = nullptr;
SetDmlOperatorDesc({ DML_OPERATOR_GEMM, &operatorDesc }, kernelCreationContext);
}
break;
case RecognizedOperatorType::ReduceSum:
{
// Get how many axes are kept in the final output, either 0 or 1 supported
// meaning full reduction or partial with one dimension left. *It could be
// generalized to support any number of output dimensions, but it would need
// to accomodate for Transposition too if the output labels are reordered.
auto keptAxes = m_components.back().GetLabels(m_labelIndices);
assert(keptAxes.size() <= 1);
// DML expects output rank to match input rank (as if ONNX ReduceSum keepdims=1).
// So replace the existing tensor description with the input sizes, except that
// reduced dimensions have size 1.
std::vector<uint32_t> reducedAxes;
std::vector<uint32_t> inputSizes = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0);
std::vector<uint32_t> outputSizes = inputSizes;
// Determine which axes are being reduced by taking the opposite of those kept.
uint32_t keptAxesMask = 0;
for (auto axis : keptAxes)
{
keptAxesMask |= (1 << axis);
}
for (uint32_t axis = 0, axisCount = static_cast<uint32_t>(outputSizes.size()); axis < axisCount; ++axis)
{
if (~keptAxesMask & (1<<axis))
{
reducedAxes.push_back(axis);
outputSizes[axis] = 1;
}
}
m_inputTensorDescs.front() = TensorDesc(m_inputTensorDescs.front().GetDmlDataType(), inputSizes, std::nullopt, 0);
m_outputTensorDescs.front() = TensorDesc(m_outputTensorDescs.front().GetDmlDataType(), outputSizes, std::nullopt, 0);
m_inputTensorDescs.front().GetDmlDesc(); // Discard value, but keep side effect of refreshing the DML view.
m_outputTensorDescs.front().GetDmlDesc(); // Discard value, but keep side effect of refreshing the DML view.
DML_REDUCE_OPERATOR_DESC operatorDesc = {};
operatorDesc.InputTensor = inputDescs.data();
operatorDesc.OutputTensor = outputDescs.data();
operatorDesc.Function = DML_REDUCE_FUNCTION_SUM;
operatorDesc.Axes = reducedAxes.data();
operatorDesc.AxisCount = gsl::narrow_cast<uint32_t>(reducedAxes.size());
SetDmlOperatorDesc({ DML_OPERATOR_REDUCE, &operatorDesc }, kernelCreationContext);
}
break;
case RecognizedOperatorType::Transpose:
case RecognizedOperatorType::Identity:
{
if (m_recognizedOperatorType == RecognizedOperatorType::Transpose)
{
// Transpose via input strides. The output tensor is not strided.
assert(m_components.front().GetDimensionCount() == m_components.back().GetDimensionCount());
auto originalStrides = m_inputTensorDescs.front().GetStrides();
std::vector<uint32_t> inputSizes = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0);
std::vector<uint32_t> inputStrides(inputSizes.size());
// If there were no strides, compute them based in descending packed order
// based on the input sizes.
if (originalStrides.empty())
{
Dml::GetDescendingPackedStrides(inputSizes, /*out*/ inputStrides);
}
else // Copy the original strides.
{
assert(originalStrides.size() >= inputStrides.size());
size_t offset = originalStrides.size() - inputStrides.size();
inputStrides.assign(originalStrides.begin() + offset, originalStrides.end());
}
// Remap transposed strides using the component labels from input to output.
auto labelIndices = m_components.back().GetLabels(m_labelIndices);
std::vector<uint32_t> newStrides(inputStrides.size());
std::vector<uint32_t> newSizes(inputStrides.size());
for (size_t i = 0, dimensionCount = inputStrides.size(); i < dimensionCount; ++i)
{
uint32_t labelIndex = labelIndices[i];
assert(labelIndex < inputStrides.size());
newSizes[i] = inputSizes[labelIndex];
newStrides[i] = inputStrides[labelIndex];
}
// Override the initial input tensor with the new strides.
m_inputTensorDescs.front() = TensorDesc(m_inputTensorDescs.front().GetDmlDataType(), newSizes, newStrides, 0);
m_outputTensorDescs.front() = TensorDesc(m_outputTensorDescs.front().GetDmlDataType(), newSizes, std::nullopt, 0);
m_inputTensorDescs.front().GetDmlDesc(); // Discard value, but keep side effect of refreshing the DML view.
m_outputTensorDescs.front().GetDmlDesc(); // Discard value, but keep side effect of refreshing the DML view.
}
DML_ELEMENT_WISE_IDENTITY_OPERATOR_DESC operatorDesc = {};
operatorDesc.InputTensor = inputDescs.data();
operatorDesc.OutputTensor = outputDescs.data();
SetDmlOperatorDesc({ DML_OPERATOR_ELEMENT_WISE_IDENTITY, &operatorDesc}, kernelCreationContext);
}
break;
default:
return;
}
}
};
void CALLBACK QueryEinSum(IMLOperatorSupportQueryContextPrivate* context, bool* isSupported)
{
*isSupported = false;
MLOperatorAttributes attributes(context);
EinSumHelper helper(attributes);
auto recognizedOperatorType = helper.GetRecognizedOperatorType();
static_assert(EinSumHelper::RecognizedOperatorType::Total == static_cast<EinSumHelper::RecognizedOperatorType>(8), "Verify this test still matches the switch above.");
*isSupported = (recognizedOperatorType != EinSumHelper::RecognizedOperatorType::None);
}
DML_OP_DEFINE_CREATION_FUNCTION(Einsum12, VersionedKernel<DmlOperatorEinSum, 12>);
} // namespace Dml

View file

@ -251,17 +251,6 @@ public:
}
};
// A specific type of operation for registration.
template <uint32_t OpsetVersion>
struct DmlOperatorResizeTemplate : public DmlOperatorResize
{
public:
DmlOperatorResizeTemplate(const MLOperatorKernelCreationContext& kernelInfo)
: DmlOperatorResize(kernelInfo, OpsetVersion)
{
}
};
void CALLBACK QueryResize(IMLOperatorSupportQueryContextPrivate* context, bool* isSupported)
{
*isSupported = false;

View file

@ -19,41 +19,25 @@ public:
ML_CHECK_VALID_ARGUMENT(kernelInfo.GetOutputCount() >= 1);
DmlOperator::Initialize(kernelInfo);
const MLOperatorEdgeDescription inputEdgeDescription = kernelInfo.GetInputEdgeDescription(0);
const std::vector<uint32_t> originalSizes = kernelInfo.GetTensorShapeDescription().GetInputTensorShape(0);
ML_CHECK_VALID_ARGUMENT(m_permutations.size() == originalSizes.size());
// Calculate strides from original shape.
ML_CHECK_VALID_ARGUMENT(!originalSizes.empty());
std::vector<uint32_t> inputStrides(originalSizes.size());
inputStrides.back() = 1;
for (int i = gsl::narrow_cast<int>(inputStrides.size()) - 2; i >= 0; i--)
{
inputStrides[i] = inputStrides[i + 1] * gsl::narrow_cast<uint32_t>(originalSizes[i + 1]);
}
Dml::GetDescendingPackedStrides(originalSizes, /*out*/ inputStrides);
const int leadingDims = gsl::narrow_cast<int32_t>(m_inputTensorDescs.front().GetDimensionCount() - originalSizes.size());
std::vector<uint32_t> sizes(m_inputTensorDescs.front().GetDimensionCount());
std::vector<uint32_t> strides(m_inputTensorDescs.front().GetDimensionCount());
// Fill leading tensor desc sizes/strides with defaults.
for (int dimDML = 0; dimDML < leadingDims; ++dimDML)
{
sizes[dimDML] = 1;
strides[dimDML] = 0;
}
std::vector<uint32_t> sizes(inputStrides.size());
std::vector<uint32_t> strides(inputStrides.size());
// Permute the shape and strides.
for (int dimInput = 0, dimCount = gsl::narrow_cast<int>(originalSizes.size()); dimInput < dimCount; ++dimInput)
{
int dimDML = dimInput + leadingDims;
int dimPermuted = m_permutations[dimInput];
ML_CHECK_VALID_ARGUMENT(gsl::narrow_cast<size_t>(dimPermuted) < originalSizes.size());
sizes[dimDML] = gsl::narrow_cast<int32_t>(originalSizes[dimPermuted]);
strides[dimDML] = inputStrides[dimPermuted];
sizes[dimInput] = originalSizes[dimPermuted];
strides[dimInput] = inputStrides[dimPermuted];
}
// Override the initial tensor descs. The output tensor is not strided.

View file

@ -143,6 +143,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(Mean);
DML_OP_EXTERN_CREATION_FUNCTION(Max);
DML_OP_EXTERN_CREATION_FUNCTION(Min);
DML_OP_EXTERN_CREATION_FUNCTION(ReduceSum);
DML_OP_EXTERN_CREATION_FUNCTION(Einsum12);
DML_OP_EXTERN_CREATION_FUNCTION(ReduceMean);
DML_OP_EXTERN_CREATION_FUNCTION(ReduceProd);
DML_OP_EXTERN_CREATION_FUNCTION(ReduceLogSum);
@ -244,6 +245,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(ConvInteger);
DML_OP_EXTERN_QUERY_FUNCTION(MaxPool);
DML_OP_EXTERN_QUERY_FUNCTION(Slice);
DML_OP_EXTERN_QUERY_FUNCTION(Resize);
DML_OP_EXTERN_QUERY_FUNCTION(EinSum);
constexpr static std::array<const char*, 1> typeNameListDefault = {"T"};
constexpr static std::array<const char*, 2> typeNameListTwo = { "T1", "T2" };
@ -460,6 +462,7 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation
{REG_INFO( 9, Where, typeNameListWhere, supportedTypeListWhere, DmlGraphSupport::Supported)},
{REG_INFO( 7, ReduceSum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO( 11, ReduceSum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_VER( 12, Einsum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryEinSum )},
{REG_INFO( 7, ReduceMean, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO( 11, ReduceMean, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO( 7, ReduceProd, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
@ -600,7 +603,7 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry)
MLOperatorKernelDescription desc = {};
desc.domain = information.domain;
desc.name = information.operatorName;
desc.executionType = MLOperatorExecutionType::D3D12;
desc.executionType = MLOperatorExecutionType::D3D12;
// 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);

View file

@ -269,7 +269,7 @@ gsl::span<const uint32_t> TensorDesc::GetStrides() const
{
return {};
}
return { m_strides, m_strides + m_bufferTensorDesc.DimensionCount };
return { m_strides, m_strides + m_bufferTensorDesc.DimensionCount };
}
DML_TENSOR_DESC TensorDesc::GetDmlDesc()

View file

@ -33,6 +33,7 @@ namespace AttrName
static constexpr const char* Dtype = "dtype";
static constexpr const char* Ends = "ends";
static constexpr const char* Epsilon = "epsilon";
static constexpr const char* Equation = "equation";
static constexpr const char* ExcludeOutside = "exclude_outside";
static constexpr const char* Exclusive = "exclusive";
static constexpr const char* Exponent = "exponent";

View file

@ -858,7 +858,264 @@ namespace OperatorHelper
std::iota(m_axes.begin(), m_axes.end(), 0);
}
}
void EinSumHelper::Initialize()
{
ParseEquationComponents();
m_recognizedOperatorType = DetermineRecognizedOperatorType();
}
void EinSumHelper::ParseEquationComponents()
{
// Parse an equation like 'ij,jk->ik' into components {ij, jk, ik} mapping letters to
// numeric indices {(0,1}, {1,2}, {0,2}}. The last component is the output.
std::map<char, uint32_t> labelMap;
std::set<char> repeatedLabels;
uint32_t currentLabelIndex = 0;
Component currentComponent = {};
bool foundOutput = false;
bool reachedEnd = false;
// Read first to last character in equation, looking for letters, commas, and one arrow.
for (char* token = m_equation.data(); !reachedEnd; ++token)
{
char ch = *token;
// Only ASCII letters are valid subscript symbols in numpy.einsum().
if ((ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z'))
{
// Check whether label already has an index.
const auto [i, inserted] = labelMap.insert({ch, currentLabelIndex});
if (inserted)
{
ML_CHECK_VALID_ARGUMENT(!foundOutput, "Found label in equation output not matching any label from inputs.")
++currentLabelIndex; // New label found.
}
else if (!foundOutput)
{
// If label in input already found earlier, then keep track of this later
// to generate the default output in case one is not specified.
repeatedLabels.insert(ch);
}
m_labelIndices.push_back(i->second);
}
else if (ch == ' ')
{
// Ignore spaces.
}
else
{
currentComponent.labelIndexEnd = static_cast<uint32_t>(m_labelIndices.size());
m_components.push_back(currentComponent);
currentComponent.labelIndexBegin = currentComponent.labelIndexEnd;
switch (ch)
{
case ',':
// Note it's valid for 2 commas be adjacent, which indicates a scalar and generates
// an empty component.
break;
case '-': // Start of "->" (must be atomic, no space between them).
++token; // Skip '-'.
ML_CHECK_VALID_ARGUMENT(*token == '>', "Expected '->' for output.")
ML_CHECK_VALID_ARGUMENT(foundOutput == false, "Only one output arrow '->' is valid.")
foundOutput = true;
break;
case '.':
// Ellipsis is unsupported. Leave recognized operator as None, deferring to another EP.
m_components.clear();
return;
case '\0':
reachedEnd = true;
break; // End of string.
default:
ML_INVALID_ARGUMENT("Unsupported character in equation string. Must be a-z, A-Z, ',', or '->'.");
}
}
}
if (!foundOutput)
{
// If no explicit output was given, generate an implicit output by ordering all the
// labels in alphabetic order (by ASCII value consistent with numpy, so Z < a).
// Exclude any labels that occurred more than once, as these cancel out.
for (auto i : labelMap)
{
if (repeatedLabels.count(i.first) == 0)
{
m_labelIndices.push_back(i.second);
}
}
// Push the final component, which is the output.
currentComponent.labelIndexEnd = static_cast<uint32_t>(m_labelIndices.size());
m_components.push_back(currentComponent);
}
}
EinSumHelper::RecognizedOperatorType EinSumHelper::DetermineRecognizedOperatorType()
{
if (m_components.empty())
{
return RecognizedOperatorType::None; // Parsing may have found unsupported components - treating as unknown.
}
// std::ranges::equal is not supported yet.
auto equals = [](gsl::span<const uint32_t> a, gsl::span<const uint32_t> b)
{
return std::equal(a.begin(), a.end(), b.begin(), b.end());
};
std::array<uint32_t, 3> componentRanks;
if (m_components.size() > componentRanks.size())
{
// No recognized operator takes more than 2 inputs and 1 output.
// EinSum itself is generic and can handle any variable number of inputs,
// but DML's operators expect fixed counts.
return RecognizedOperatorType::None;
}
else if (m_components.size() == 2)
{
auto& inputLabels = m_components[0].GetLabels(m_labelIndices);
auto& outputLabels = m_components[1].GetLabels(m_labelIndices);
if (inputLabels.size() == outputLabels.size())
{
// Check identity.
if (equals(inputLabels, outputLabels))
{
// Handles: "->", "i->i", "ij->ij", "ijk->ijk", "ijkl->ijkl" ...
return RecognizedOperatorType::Identity;
}
else // Transpose since a permutation exists.
{
// Handles: "ij->ji", "ijk->kji", "ijkl->lkji", "ijkl->ijkl" ...
return RecognizedOperatorType::Transpose;
}
}
else if (outputLabels.empty()) // Scalar output, with all inputs reduced.
{
// Handles: "i->", "ij->", "ijk->", "ijkl->" ...
return RecognizedOperatorType::ReduceSum;
}
}
else if (m_components.size() == 3)
{
// If all components have the same size and label order, then apply elementwise multiplication.
auto& inputALabels = m_components[0].GetLabels(m_labelIndices);
auto& inputBLabels = m_components[1].GetLabels(m_labelIndices);
auto& outputLabels = m_components[2].GetLabels(m_labelIndices);
if (equals(inputALabels, outputLabels) && equals(inputBLabels, outputLabels))
{
// Handles: "i,i->i", "ij,ij->ij", "ijk,ijk->ijk", "ijkl,ijkl->ijkl" ...
return RecognizedOperatorType::Multiply;
}
}
// Otherwise check for special cases of dedicated operators...
struct RecognizedOperatorInfo
{
RecognizedOperatorType recognizedOperatorType;
std::initializer_list<const uint32_t> componentRanks;
std::initializer_list<const uint32_t> labelIndices;
};
const RecognizedOperatorInfo recognizedOperators[] = {
{RecognizedOperatorType::MatMul, {2,2,2},{0,1, 1,2, 0,2}}, // ij,jk->ik
{RecognizedOperatorType::MatMul, {3,3,3},{0,1,2, 0,2,3, 0,1,3}}, // bij,bjk->bik
{RecognizedOperatorType::MatMul, {4,4,4},{0,1,2,3, 0,1,3,4, 0,1,2,4}}, // abij,abjk->abik
{RecognizedOperatorType::MatMulTransposeA, {2,2,2},{0,1, 0,2, 1,2}}, // ji,jk->ik
{RecognizedOperatorType::MatMulTransposeA, {3,3,3},{0,1,2, 0,1,3, 0,2,3}}, // bji,bjk->bik
{RecognizedOperatorType::MatMulTransposeA, {4,4,4},{0,1,2,3, 0,1,2,4, 0,1,3,4}}, // abji,abjk->abik
{RecognizedOperatorType::MatMulTransposeB, {2,2,2},{0,1, 2,1, 0,2}}, // ij,kj->ik
{RecognizedOperatorType::MatMulTransposeB, {3,3,3},{0,1,2, 0,3,2, 0,1,3}}, // bij,bkj->bik
{RecognizedOperatorType::MatMulTransposeB, {4,4,4},{0,1,2,3, 0,1,4,3, 0,1,2,4}}, // abij,abkj->abik
{RecognizedOperatorType::MatMulTransposeB, {1,1,0},{0,0,}}, // i,i-> (1D inner_prod)
{RecognizedOperatorType::ReduceSum, {2,1 },{0,1, 0}}, // ij->i
{RecognizedOperatorType::ReduceSum, {2,1 },{0,1, 1}}, // ij->j
};
// For each recognized operator, compare the labels-per-component and label indices.
for (auto& recognizedOperator : recognizedOperators)
{
if (equals(m_labelIndices, recognizedOperator.labelIndices)
&& m_components.size() == recognizedOperator.componentRanks.size())
{
for (size_t i = 0; i < m_components.size(); ++i)
{
componentRanks[i] = m_components[i].GetDimensionCount();
}
if (equals(gsl::make_span(componentRanks.data(), m_components.size()), recognizedOperator.componentRanks))
{
return recognizedOperator.recognizedOperatorType;
}
}
}
return RecognizedOperatorType::None;
}
std::vector<EdgeShapes> EinSumHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
assert(!m_components.empty()); // Should have already parsed components.
uint32_t inputCount = shapeInfo.GetInputCount();
uint32_t outputCount = shapeInfo.GetOutputCount();
ML_CHECK_VALID_ARGUMENT(inputCount + 1 == m_components.size(), "Mismatch between input tensor count and string equation component count.");
ML_CHECK_VALID_ARGUMENT(outputCount == 1, "EinSum expects exactly 1 output tensor.");
std::vector<uint32_t> labelSizes(m_labelIndices.size(), INT_MIN);
// Read every input tensor, comparing labels to ensure consistent sizes from the equation parsed earlier.
for (uint32_t i = 0; i < inputCount; ++i)
{
auto inputShape = shapeInfo.GetInputTensorShape(i);
auto& component = m_components[i];
auto labelIndices = component.GetLabels(m_labelIndices);
uint32_t dimensionCount = component.GetDimensionCount();
ML_CHECK_VALID_ARGUMENT(inputShape.size() == dimensionCount, "Mismatch between input tensor shape and string equation label count.");
for (uint32_t i = 0; i < dimensionCount; ++i)
{
// If this is the first time seeing this label, then record the size.
// Otherwise any following occurrences of the label must match sizes.
// e.g. Given "ij,ji", both i's and both j's must match dimension sizes.
uint32_t dimensionSize = inputShape[i];
uint32_t labelIndex = labelIndices[i];
assert(labelIndex < labelSizes.size());
if (labelSizes[labelIndex] == INT_MIN)
{
labelSizes[labelIndex] = dimensionSize;
}
else
{
ML_CHECK_VALID_ARGUMENT(labelSizes[labelIndex] == dimensionSize, "All labels must have the same dimension sizes.");
}
}
}
// Generate output dimensions from corresponding input tensor labels.
// e.g. Given ij,jk->ij with [2,3] and [3,5], the output is [2,5].
std::vector<uint32_t> outputDimensions;
auto outputLabelIndices = m_components.back().GetLabels(m_labelIndices);
for (auto labelIndex : outputLabelIndices)
{
outputDimensions.push_back(labelSizes[labelIndex]);
}
return { std::move(EdgeShapes(outputDimensions)) };
}
std::vector<EdgeShapes> MatMulHelperBase::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
ML_CHECK_VALID_ARGUMENT(shapeInfo.GetInputCount() >= 2);

View file

@ -771,6 +771,70 @@ class ReduceHelper : public ReduceHelperBase {
ReduceHelper(const Info_t& info, const Shape_t& shape) : ReduceHelperBase(info, shape, true) {}
};
class EinSumHelper
{
public:
void Initialize();
// Info_t is used to obtain attributes which will be used for calculating the output shape later.
// Shape_t is used to obtain input shape which will be used for adjusting attribute value.
template <typename Info_t, typename Shape_t>
EinSumHelper(const Info_t& info, const Shape_t& shape, uint32_t opsetVersion)
{
m_equation = info.GetAttribute(AttrName::Equation);
Initialize();
}
EinSumHelper(const MLOperatorAttributes& info)
{
m_equation = info.GetAttribute(AttrName::Equation);
Initialize();
}
std::vector<EdgeShapes> GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const;
enum class RecognizedOperatorType
{
None,
Identity,
Multiply,
MatMul,
MatMulTransposeA,
MatMulTransposeB,
ReduceSum,
Transpose,
Total,
};
RecognizedOperatorType GetRecognizedOperatorType() const noexcept { return m_recognizedOperatorType; }
protected:
void ParseEquationComponents();
RecognizedOperatorType DetermineRecognizedOperatorType();
protected:
struct Component
{
uint32_t labelIndexBegin;
uint32_t labelIndexEnd;
uint32_t GetDimensionCount() const noexcept
{
return labelIndexEnd - labelIndexBegin;
}
gsl::span<const uint32_t> GetLabels(gsl::span<const uint32_t> labels) const
{
return labels.subspan(labelIndexBegin, labelIndexEnd - labelIndexBegin);
};
};
std::string m_equation;
std::vector<uint32_t> m_labelIndices; // Concatenation of all labels as rebased indices ("ij,ai" -> 0,1,2,0).
std::vector<Component> m_components; // All components in order, including inputs and output.
std::vector<uint32_t> m_outputDimensions;
RecognizedOperatorType m_recognizedOperatorType = RecognizedOperatorType::None;
};
class MatMulHelperBase {
public:
// Info_t is used to obtain attributes which will be used for calculating the output shape later.
@ -1465,6 +1529,7 @@ using ShapeInferenceHelper_ReduceL1 = ReduceHelper;
using ShapeInferenceHelper_ReduceL2 = ReduceHelper;
using ShapeInferenceHelper_ReduceMax = ReduceHelper;
using ShapeInferenceHelper_ReduceMin = ReduceHelper;
using ShapeInferenceHelper_Einsum12 = VersionedOpsetHelper<EinSumHelper, 12>;
using ShapeInferenceHelper_ArgMax = ArgMinArgMaxHelper;
using ShapeInferenceHelper_ArgMin = ArgMinArgMaxHelper;
using ShapeInferenceHelper_Gemm = GemmHelper;

View file

@ -247,19 +247,20 @@ namespace OperatorHelper
namespace OnnxOperatorSet12
{
static const int sc_sinceVer_GreaterOrEqual = 12;
static const int sc_sinceVer_LessOrEqual = 12;
static const int sc_sinceVer_ArgMin = 12;
static const int sc_sinceVer_ArgMax = 12;
static const int sc_sinceVer_Celu = 12;
static const int sc_sinceVer_Clip = 12;
static const int sc_sinceVer_Einsum = 12;
static const int sc_sinceVer_GatherND = 12;
static const int sc_sinceVer_GreaterOrEqual = 12;
static const int sc_sinceVer_LessOrEqual = 12;
static const int sc_sinceVer_MaxPool = 12;
static const int sc_sinceVer_Min = 12;
static const int sc_sinceVer_Max = 12;
static const int sc_sinceVer_Pow = 12;
static const int sc_sinceVer_MaxPool = 12;
static const int sc_sinceVer_ReduceMax = 12;
static const int sc_sinceVer_ReduceMin = 12;
static const int sc_sinceVer_ArgMin = 12;
static const int sc_sinceVer_ArgMax = 12;
} // namespace OnnxOperatorSet12
namespace MsftOperatorSet1

View file

@ -7,6 +7,8 @@
#include <cassert>
#include <chrono>
#include <vector>
#include <map>
#include <set>
#include <numeric>
#include <wrl/client.h>