diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md
index c4984e423d..7802a6853a 100644
--- a/docs/OperatorKernels.md
+++ b/docs/OperatorKernels.md
@@ -1253,6 +1253,7 @@ Do not modify directly.*
|QLinearSigmoid|*in* X:**T**
*in* X_scale:**tensor(float)**
*in* X_zero_point:**T**
*in* Y_scale:**tensor(float)**
*in* Y_zero_point:**T**
*out* Y:**T**|1+|**T** = tensor(int8), tensor(uint8)|
|QuantizeLinear|*in* x:**T1**
*in* y_scale:**T1**
*in* y_zero_point:**T2**
*out* y:**T2**|1+|**T1** = tensor(float), tensor(float16), tensor(int32)
**T2** = tensor(int8), tensor(uint8)|
|QuickGelu|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
+|RotaryEmbedding|*in* input:**T**
*in* position_ids:**M**
*in* cos_cache:**T**
*in* sin_cache:**T**
*out* output:**T**|1+|**M** = tensor(int64)
**T** = tensor(float), tensor(float16)|
|SkipLayerNormalization|*in* input:**T**
*in* skip:**T**
*in* gamma:**T**
*in* beta:**T**
*in* bias:**T**
*out* output:**T**
*out* mean:**U**
*out* inv_std_var:**U**
*out* input_skip_bias_sum:**T**|1+|**T** = tensor(float), tensor(float16)|
| |
| |
diff --git a/onnxruntime/contrib_ops/cuda/bert/skip_layer_norm_impl.cu b/onnxruntime/contrib_ops/cuda/bert/skip_layer_norm_impl.cu
index 973ef8d304..50c8e4b5e0 100644
--- a/onnxruntime/contrib_ops/cuda/bert/skip_layer_norm_impl.cu
+++ b/onnxruntime/contrib_ops/cuda/bert/skip_layer_norm_impl.cu
@@ -51,7 +51,7 @@ half maybe2half(float x) {
// Using only power of 2 numbers will lead to waste of compute for same size such as 768, which is a very common case
// in BERT. Ideally we can step by wrap_size * num_unroll, but listing too many steps will cause long compile time.
-constexpr int kSizes[] = {128, 384, 768, 1024, 2048, 4096, 5120, 8192};
+constexpr int kSizes[] = {128, 320, 384, 640, 768, 1024, 1280, 2048, 4096, 5120, 8192};
constexpr size_t kNumOfSizes = sizeof(kSizes) / sizeof(kSizes[0]);
constexpr int kMaxSize = kSizes[kNumOfSizes - 1];
constexpr int kMinBlockSize = 32;
@@ -206,7 +206,7 @@ void LaunchSkipLayerNormKernel(
#define CASE_NEXT_SIZE(next_size_value) \
case next_size_value: { \
static_assert(next_size_value >= kSizes[0] && next_size_value <= kMaxSize); \
- if constexpr (next_size_value >= 8 * 256) { \
+ if constexpr (next_size_value >= 320) { \
if (can_unroll_vec8) { \
constexpr int block_size = next_size_value / 8; \
LAUNCH_SKIP_LAYER_NORM_KERNEL_SMALL(8); \
@@ -239,6 +239,9 @@ void LaunchSkipLayerNormKernel(
CASE_NEXT_SIZE(kSizes[5]);
CASE_NEXT_SIZE(kSizes[6]);
CASE_NEXT_SIZE(kSizes[7]);
+ CASE_NEXT_SIZE(kSizes[8]);
+ CASE_NEXT_SIZE(kSizes[9]);
+ CASE_NEXT_SIZE(kSizes[10]);
default: {
constexpr int block_size = 256;
LAUNCH_SKIP_LAYER_NORM_KERNEL();
diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRotaryEmbedding.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRotaryEmbedding.cpp
new file mode 100644
index 0000000000..30c339b845
--- /dev/null
+++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRotaryEmbedding.cpp
@@ -0,0 +1,436 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+#include "precomp.h"
+
+// This operator is easier to understand by looking at a python implementation of the non-interleaved version:
+//
+// def rotate_half(x):
+// """Rotates half the hidden dims of the input."""
+// half_dim = x.shape[-1] // 2
+// x1 = x[..., :half_dim]
+// x2 = x[..., half_dim:]
+// return np.concatenate((-x2, x1), dim=-1)
+//
+//
+// def apply_rope(x, cos, sin, position_ids):
+// cos = cos[position_ids].unsqueeze(1) # [bs, 1, seq_len, dim]
+// sin = sin[position_ids].unsqueeze(1) # [bs, 1, seq_len, dim]
+// x_embed = (x * cos) + (rotate_half(x) * sin)
+// return x_embed
+//
+// For the non-interleaved version, we multiply the cos cache by the non-rotated input tensor while we multiply the sin cache
+// by the rotated input tensor. Rotating the tensor means slicing it in half on the head dimension and swapping the 2 halves.
+//
+// The interleaved version is very similar but instead of swapping 2 halves, we swap every pair of adjacent elements and we swap
+// the sign of every adjacent element.
+
+namespace Dml
+{
+class DmlOperatorRotaryEmbedding : public DmlOperator
+{
+public:
+ DmlOperatorRotaryEmbedding(const MLOperatorKernelCreationContext& kernelInfo) : DmlOperator(kernelInfo)
+ {
+ enum InputIndex : uint32_t
+ {
+ inputDataIndex,
+ positionIdsIndex,
+ cosCacheIndex,
+ sinCacheIndex,
+ };
+
+ ML_CHECK_VALID_ARGUMENT(kernelInfo.GetInputCount() == 4);
+ ML_CHECK_VALID_ARGUMENT(kernelInfo.GetOutputCount() == 1);
+
+ // When positionIds is a scalar, it represents the start offset for each sequence
+ const bool positionIdsIsOffset = kernelInfo.GetInputTensorDimensionCount(positionIdsIndex) == 1;
+
+ Initialize(kernelInfo);
+
+ ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[inputDataIndex].GetDimensionCount() == 4);
+ ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[positionIdsIndex].GetDimensionCount() == 4);
+ ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[cosCacheIndex].GetDimensionCount() == 4);
+ ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[sinCacheIndex].GetDimensionCount() == 4);
+
+ ML_CHECK_VALID_ARGUMENT(m_outputTensorDescs[0].GetDimensionCount() == 4);
+
+ ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[cosCacheIndex].GetSizes() == m_inputTensorDescs[sinCacheIndex].GetSizes());
+ const uint32_t headSize = m_inputTensorDescs[cosCacheIndex].GetSizes().back() * 2;
+
+ // The last dimension of the data is the hidden size, so it must be divisible by the head size
+ ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[inputDataIndex].GetSizes().back() % headSize == 0);
+
+ // We resize the data to be of shape [batchSize, sequenceLength, numHeads, headSize]
+ const auto inputDataSizes = m_inputTensorDescs[inputDataIndex].GetSizes();
+ const uint32_t batchSize = inputDataSizes[1];
+ const uint32_t sequenceLength = inputDataSizes[2];
+ const uint32_t numHeads = inputDataSizes[3] / headSize;
+
+ const auto cosCacheSizes = m_inputTensorDescs[cosCacheIndex].GetSizes();
+ const uint32_t maxSequenceLength = cosCacheSizes[cosCacheSizes.size() - 2];
+
+ if (sequenceLength > maxSequenceLength)
+ {
+ ORT_NOT_IMPLEMENTED("Updating cos_cache and sin_cache in RotaryEmbedding is not currently supported");
+ }
+
+ const bool interleaved = gsl::narrow_cast(kernelInfo.GetOptionalAttribute(AttrName::Interleaved, 0));
+
+ std::vector inputDescs = GetDmlInputDescs();
+ const MLOperatorTensorDataType dataType = kernelInfo.GetInputEdgeDescription(inputDataIndex).tensorDataType;
+
+ // Splitting the hiddenSize into numHeads and headSize dimensions makes it easier for DML to handle
+ const std::array inputOutputShape = {batchSize, sequenceLength, numHeads, headSize};
+ TensorDesc inputOutputTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, inputOutputShape);
+ const DML_TENSOR_DESC inputOutputDmlTensorDesc = inputOutputTensorDesc.GetDmlDesc();
+
+ // Copy the input to preserve its real input shape in the graph without reshaping it. This will disappear during DML's graph compilation phase.
+ DML_SCALE_BIAS scaleBias = {1.0f, 0.0f};
+
+ DML_ELEMENT_WISE_IDENTITY_OPERATOR_DESC copyInputDesc{};
+ copyInputDesc.InputTensor = &inputOutputDmlTensorDesc;
+ copyInputDesc.OutputTensor = &inputOutputDmlTensorDesc;
+ copyInputDesc.ScaleBias = &scaleBias;
+ const DML_OPERATOR_DESC copyInputDmlDesc = {DML_OPERATOR_ELEMENT_WISE_IDENTITY, ©InputDesc};
+
+ // Split the input data into 2 equal parts
+ const std::vector inputDataTensorShape = interleaved
+ ? std::vector({batchSize, sequenceLength, numHeads, headSize / 2, 2})
+ : std::vector({batchSize, sequenceLength, numHeads, 2, headSize / 2});
+
+ const std::vector splitInputDataTensorShape = interleaved
+ ? std::vector({batchSize, sequenceLength, numHeads, headSize / 2, 1})
+ : std::vector({batchSize, sequenceLength, numHeads, 1, headSize / 2});
+
+ TensorDesc inputDataTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, inputDataTensorShape);
+ const DML_TENSOR_DESC inputDataDmlTensorDesc = inputDataTensorDesc.GetDmlDesc();
+
+ TensorDesc splitInputDataTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, splitInputDataTensorShape);
+ const std::array splitInputDataDmlTensorDescs = {splitInputDataTensorDesc.GetDmlDesc(), splitInputDataTensorDesc.GetDmlDesc()};
+
+ DML_SPLIT_OPERATOR_DESC splitInputDesc{};
+ splitInputDesc.InputTensor = &inputDataDmlTensorDesc;
+ splitInputDesc.OutputTensors = splitInputDataDmlTensorDescs.data();
+ splitInputDesc.OutputCount = gsl::narrow_cast(splitInputDataDmlTensorDescs.size());
+ splitInputDesc.Axis = interleaved
+ ? gsl::narrow_cast(splitInputDataTensorShape.size()) - 1
+ : gsl::narrow_cast(splitInputDataTensorShape.size()) - 2;
+
+ const DML_OPERATOR_DESC splitInputDmlDesc = {DML_OPERATOR_SPLIT, &splitInputDesc};
+
+ // Swap the 2 halves and join them together
+ DML_JOIN_OPERATOR_DESC joinInputDesc{};
+ joinInputDesc.InputTensors = splitInputDataDmlTensorDescs.data();
+ joinInputDesc.OutputTensor = &inputDataDmlTensorDesc;
+ joinInputDesc.Axis = splitInputDesc.Axis;
+ joinInputDesc.InputCount = gsl::narrow_cast(splitInputDataDmlTensorDescs.size());
+ const DML_OPERATOR_DESC joinInputDmlDesc = {DML_OPERATOR_JOIN, &joinInputDesc};
+
+ // We generate a sequence from 0 to sequenceLength and add the offset to it
+ const std::array positionIdsRangeShape = {1, 1, 1, sequenceLength};
+ auto positionIdsDataType = kernelInfo.GetInputEdgeDescription(positionIdsIndex).tensorDataType;
+ TensorDesc positionIdsRangeTensorDesc = TensorDesc::ConstructDefaultTensorDesc(positionIdsDataType, positionIdsRangeShape);
+ const DML_TENSOR_DESC positionIdsRangeDmlTensorDesc = positionIdsRangeTensorDesc.GetDmlDesc();
+
+ const std::array broadcastedPositionIdsRangeShape = {1, 1, batchSize, sequenceLength};
+ TensorDesc broadcastedPositionIdsRangeTensorDesc = TensorDesc::ConstructBroadcastedTensorDesc(positionIdsDataType, broadcastedPositionIdsRangeShape, positionIdsRangeShape);
+ const DML_TENSOR_DESC broadcastedPositionIdsRangeDmlTensorDesc = broadcastedPositionIdsRangeTensorDesc.GetDmlDesc();
+
+ const std::array broadcastedOffsetShape = {1, 1, batchSize, sequenceLength};
+ TensorDesc broadcastedOffsetTensorDesc = TensorDesc::ConstructBroadcastedTensorDesc(positionIdsDataType, broadcastedOffsetShape, m_inputTensorDescs[positionIdsIndex].GetSizes());
+ const DML_TENSOR_DESC broadcastedOffsetDmlTensorDesc = broadcastedOffsetTensorDesc.GetDmlDesc();
+
+ TensorDesc offsetPositionIdsTensorDesc = TensorDesc::ConstructDefaultTensorDesc(positionIdsDataType, broadcastedOffsetShape);
+ const DML_TENSOR_DESC offsetPositionIdsRangeDmlTensorDesc = offsetPositionIdsTensorDesc.GetDmlDesc();
+
+ DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC positionIdsRange{};
+ DML_ELEMENT_WISE_ADD_OPERATOR_DESC positionIdsAddOffset{};
+ if (positionIdsIsOffset)
+ {
+ ML_CHECK_VALID_ARGUMENT(positionIdsDataType == MLOperatorTensorDataType::Int64);
+ positionIdsRange.ValueDataType = DML_TENSOR_DATA_TYPE_INT64;
+ positionIdsRange.ValueDelta.Int64 = 1;
+ positionIdsRange.OutputTensor = &positionIdsRangeDmlTensorDesc;
+
+ positionIdsAddOffset.ATensor = &broadcastedPositionIdsRangeDmlTensorDesc;
+ positionIdsAddOffset.BTensor = &broadcastedOffsetDmlTensorDesc;
+ positionIdsAddOffset.OutputTensor = &offsetPositionIdsRangeDmlTensorDesc;
+ }
+ const DML_OPERATOR_DESC positionIdsRangeDmlDesc = {DML_OPERATOR_FILL_VALUE_SEQUENCE, &positionIdsRange};
+ const DML_OPERATOR_DESC positionIdsAddOffsetDmlDesc = {DML_OPERATOR_ELEMENT_WISE_ADD, &positionIdsAddOffset};
+
+ // Gather the cos/sin values based on the position ids
+ const std::array gatheredCosSinShape = {1, batchSize, sequenceLength, headSize / 2};
+ TensorDesc gatheredCosSinTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, gatheredCosSinShape);
+ const DML_TENSOR_DESC gatheredCosSinDmlTensorDesc = gatheredCosSinTensorDesc.GetDmlDesc();
+
+ DML_GATHER_OPERATOR_DESC gatherCosSinDesc{};
+ gatherCosSinDesc.InputTensor = &inputDescs[cosCacheIndex];
+ gatherCosSinDesc.IndicesTensor = positionIdsIsOffset ? &offsetPositionIdsRangeDmlTensorDesc : &inputDescs[positionIdsIndex];
+ gatherCosSinDesc.OutputTensor = &gatheredCosSinDmlTensorDesc;
+ gatherCosSinDesc.Axis = 2;
+ gatherCosSinDesc.IndexDimensions = 2;
+ const DML_OPERATOR_DESC gatherCosSinDmlDesc {DML_OPERATOR_GATHER, &gatherCosSinDesc};
+
+ // After gathering cos/sin, reshape and broadcast them to match the number of heads of the input data
+ const std::vector reshapedCosSinShape = interleaved
+ ? std::vector({batchSize, sequenceLength, 1, headSize / 2, 1})
+ : std::vector({batchSize, sequenceLength, 1, 1, headSize / 2});
+ TensorDesc broadcastedCosSinTensorDesc = TensorDesc::ConstructBroadcastedTensorDesc(dataType, inputDataTensorShape, reshapedCosSinShape);
+ const DML_TENSOR_DESC broadcastedCosSinDmlTensorDesc = broadcastedCosSinTensorDesc.GetDmlDesc();
+
+ // Create a vector that contains the sign values {-1, 1}
+ const std::array signTensorShape = {2};
+ TensorDesc signTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, signTensorShape);
+ const DML_TENSOR_DESC signDmlTensorDesc = signTensorDesc.GetDmlDesc();
+
+ DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC signRange{};
+ signRange.OutputTensor = &signDmlTensorDesc;
+ if (dataType == MLOperatorTensorDataType::Float16)
+ {
+ const auto valueStart = static_cast(-1.0f);
+ const auto valueDelta = static_cast(2.0f);
+ memcpy(signRange.ValueStart.Bytes, reinterpret_cast(&valueStart), sizeof(valueStart));
+ memcpy(signRange.ValueDelta.Bytes, reinterpret_cast(&valueDelta), sizeof(valueDelta));
+ signRange.ValueDataType = DML_TENSOR_DATA_TYPE_FLOAT16;
+ }
+ else
+ {
+ ML_CHECK_VALID_ARGUMENT(dataType == MLOperatorTensorDataType::Float);
+ signRange.ValueStart.Float32 = -1.0f;
+ signRange.ValueDelta.Float32 = 2.0f;
+ signRange.ValueDataType = DML_TENSOR_DATA_TYPE_FLOAT32;
+ }
+ const DML_OPERATOR_DESC signRangeDmlDesc = {DML_OPERATOR_FILL_VALUE_SEQUENCE, &signRange};
+
+ // Multiply the broadcasted sign values with the rotated input
+ const std::vector reshapedSignShape = interleaved
+ ? std::vector({1, 1, 1, 1, 2})
+ : std::vector({1, 1, 1, 2, 1});
+ TensorDesc broadcastedSignCosSinTensorDesc = TensorDesc::ConstructBroadcastedTensorDesc(dataType, inputDataTensorShape, reshapedSignShape);
+ const DML_TENSOR_DESC broadcastedSignDmlTensorDesc = broadcastedSignCosSinTensorDesc.GetDmlDesc();
+
+ DML_ELEMENT_WISE_MULTIPLY_OPERATOR_DESC mulSignDesc{};
+ mulSignDesc.ATensor = &inputDataDmlTensorDesc;
+ mulSignDesc.BTensor = &broadcastedSignDmlTensorDesc;
+ mulSignDesc.OutputTensor = &inputDataDmlTensorDesc;
+ const DML_OPERATOR_DESC mulSignDmlDesc = {DML_OPERATOR_ELEMENT_WISE_MULTIPLY, &mulSignDesc};
+
+ // Multiply the non-rotated data with the cos and the rotated data with the sin
+ DML_ELEMENT_WISE_MULTIPLY_OPERATOR_DESC mulCosSinDesc{};
+ mulCosSinDesc.ATensor = &inputDataDmlTensorDesc;
+ mulCosSinDesc.BTensor = &broadcastedCosSinDmlTensorDesc;
+ mulCosSinDesc.OutputTensor = &inputDataDmlTensorDesc;
+ const DML_OPERATOR_DESC mulCosSinDmlDesc = {DML_OPERATOR_ELEMENT_WISE_MULTIPLY, &mulCosSinDesc};
+
+ // Add the multiplied cos and sin values together
+ DML_ELEMENT_WISE_ADD_OPERATOR_DESC addDesc{};
+ addDesc.ATensor = &inputOutputDmlTensorDesc;
+ addDesc.BTensor = &inputOutputDmlTensorDesc;
+ addDesc.OutputTensor = &inputOutputDmlTensorDesc;
+ const DML_OPERATOR_DESC addDmlDesc = {DML_OPERATOR_ELEMENT_WISE_ADD, &addDesc};
+
+ // Construct the graph
+ std::vector inputEdges;
+ std::vector intermediateEdges;
+ std::vector outputEdges;
+
+ std::vector opDescs = {
+ ©InputDmlDesc, // Copy the input data to preseve the real input shape
+ &splitInputDmlDesc, // Split the input data
+ &gatherCosSinDmlDesc, // Gather cos
+ &gatherCosSinDmlDesc, // Gather sin
+ &signRangeDmlDesc, // Generate the signs
+
+ &joinInputDmlDesc, // Join the split data
+ &mulCosSinDmlDesc, // Multiply cos with the non-rotated data
+ &mulCosSinDmlDesc, // Multiply sin with the rotated data
+ &mulSignDmlDesc, // Multiply the sign with the rotated data
+ &addDmlDesc, // Add the rotated cos and non-rotated sin parts together
+ };
+
+ enum NodeIndex : uint32_t
+ {
+ copyInputOpIndex,
+ splitInputOpIndex,
+ gatherCosOpIndex,
+ gatherSinOpIndex,
+ signRangeOpIndex,
+
+ joinInputOpIndex,
+ mulCosOpIndex,
+ mulSinOpIndex,
+ mulSignOpIndex,
+ addOpIndex,
+
+ // The following indices are optional
+ positionIdsRangeOpIndex,
+ positionIdsAddOffsetOpIndex,
+ };
+
+ if (positionIdsIsOffset)
+ {
+ opDescs.push_back(&positionIdsRangeDmlDesc);
+ opDescs.push_back(&positionIdsAddOffsetDmlDesc);
+
+ DML_INPUT_GRAPH_EDGE_DESC positionIdsToAddOffsetEdge = {};
+ positionIdsToAddOffsetEdge.GraphInputIndex = positionIdsIndex;
+ positionIdsToAddOffsetEdge.ToNodeIndex = positionIdsAddOffsetOpIndex;
+ positionIdsToAddOffsetEdge.ToNodeInputIndex = 1;
+ inputEdges.push_back(positionIdsToAddOffsetEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC positionIdsOffsetToAddOffsetEdge = {};
+ positionIdsOffsetToAddOffsetEdge.FromNodeIndex = positionIdsRangeOpIndex;
+ positionIdsOffsetToAddOffsetEdge.FromNodeOutputIndex = 0;
+ positionIdsOffsetToAddOffsetEdge.ToNodeIndex = positionIdsAddOffsetOpIndex;
+ positionIdsOffsetToAddOffsetEdge.ToNodeInputIndex = 0;
+ intermediateEdges.push_back(positionIdsOffsetToAddOffsetEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC positionIdsAddOffsetToGatherCosEdge = {};
+ positionIdsAddOffsetToGatherCosEdge.FromNodeIndex = positionIdsAddOffsetOpIndex;
+ positionIdsAddOffsetToGatherCosEdge.FromNodeOutputIndex = 0;
+ positionIdsAddOffsetToGatherCosEdge.ToNodeIndex = gatherCosOpIndex;
+ positionIdsAddOffsetToGatherCosEdge.ToNodeInputIndex = 1;
+ intermediateEdges.push_back(positionIdsAddOffsetToGatherCosEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC positionIdsAddOffsetToGatherSinEdge = {};
+ positionIdsAddOffsetToGatherSinEdge.FromNodeIndex = positionIdsAddOffsetOpIndex;
+ positionIdsAddOffsetToGatherSinEdge.FromNodeOutputIndex = 0;
+ positionIdsAddOffsetToGatherSinEdge.ToNodeIndex = gatherSinOpIndex;
+ positionIdsAddOffsetToGatherSinEdge.ToNodeInputIndex = 1;
+ intermediateEdges.push_back(positionIdsAddOffsetToGatherSinEdge);
+ }
+ else
+ {
+ DML_INPUT_GRAPH_EDGE_DESC positionIdsToGatherCosEdge = {};
+ positionIdsToGatherCosEdge.GraphInputIndex = positionIdsIndex;
+ positionIdsToGatherCosEdge.ToNodeIndex = gatherCosOpIndex;
+ positionIdsToGatherCosEdge.ToNodeInputIndex = 1;
+ inputEdges.push_back(positionIdsToGatherCosEdge);
+
+ DML_INPUT_GRAPH_EDGE_DESC positionIdsToGatherSinEdge = {};
+ positionIdsToGatherSinEdge.GraphInputIndex = positionIdsIndex;
+ positionIdsToGatherSinEdge.ToNodeIndex = gatherSinOpIndex;
+ positionIdsToGatherSinEdge.ToNodeInputIndex = 1;
+ inputEdges.push_back(positionIdsToGatherSinEdge);
+ }
+
+ DML_INPUT_GRAPH_EDGE_DESC inputToCopyInputEdge = {};
+ inputToCopyInputEdge.GraphInputIndex = inputDataIndex;
+ inputToCopyInputEdge.ToNodeIndex = copyInputOpIndex;
+ inputToCopyInputEdge.ToNodeInputIndex = 0;
+ inputEdges.push_back(inputToCopyInputEdge);
+
+ DML_INPUT_GRAPH_EDGE_DESC cosToGatherEdge = {};
+ cosToGatherEdge.GraphInputIndex = cosCacheIndex;
+ cosToGatherEdge.ToNodeIndex = gatherCosOpIndex;
+ cosToGatherEdge.ToNodeInputIndex = 0;
+ inputEdges.push_back(cosToGatherEdge);
+
+ DML_INPUT_GRAPH_EDGE_DESC sinToGatherEdge = {};
+ sinToGatherEdge.GraphInputIndex = sinCacheIndex;
+ sinToGatherEdge.ToNodeIndex = gatherSinOpIndex;
+ sinToGatherEdge.ToNodeInputIndex = 0;
+ inputEdges.push_back(sinToGatherEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC inputToSplitEdge = {};
+ inputToSplitEdge.FromNodeIndex = copyInputOpIndex;
+ inputToSplitEdge.FromNodeOutputIndex = 0;
+ inputToSplitEdge.ToNodeIndex = splitInputOpIndex;
+ inputToSplitEdge.ToNodeInputIndex = 0;
+ intermediateEdges.push_back(inputToSplitEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC nonRotatedDataToMulEdge = {};
+ nonRotatedDataToMulEdge.FromNodeIndex = copyInputOpIndex;
+ nonRotatedDataToMulEdge.FromNodeOutputIndex = 0;
+ nonRotatedDataToMulEdge.ToNodeIndex = mulCosOpIndex;
+ nonRotatedDataToMulEdge.ToNodeInputIndex = 0;
+ intermediateEdges.push_back(nonRotatedDataToMulEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC secondHalfDataToJoinEdge = {};
+ secondHalfDataToJoinEdge.FromNodeIndex = splitInputOpIndex;
+ secondHalfDataToJoinEdge.FromNodeOutputIndex = 1;
+ secondHalfDataToJoinEdge.ToNodeIndex = joinInputOpIndex;
+ secondHalfDataToJoinEdge.ToNodeInputIndex = 0;
+ intermediateEdges.push_back(secondHalfDataToJoinEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC firstHalfDataToJoinEdge = {};
+ firstHalfDataToJoinEdge.FromNodeIndex = splitInputOpIndex;
+ firstHalfDataToJoinEdge.FromNodeOutputIndex = 0;
+ firstHalfDataToJoinEdge.ToNodeIndex = joinInputOpIndex;
+ firstHalfDataToJoinEdge.ToNodeInputIndex = 1;
+ intermediateEdges.push_back(firstHalfDataToJoinEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC cosToMulEdge = {};
+ cosToMulEdge.FromNodeIndex = gatherCosOpIndex;
+ cosToMulEdge.FromNodeOutputIndex = 0;
+ cosToMulEdge.ToNodeIndex = mulCosOpIndex;
+ cosToMulEdge.ToNodeInputIndex = 1;
+ intermediateEdges.push_back(cosToMulEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC rotatedDataToMulEdge = {};
+ rotatedDataToMulEdge.FromNodeIndex = joinInputOpIndex;
+ rotatedDataToMulEdge.FromNodeOutputIndex = 0;
+ rotatedDataToMulEdge.ToNodeIndex = mulSinOpIndex;
+ rotatedDataToMulEdge.ToNodeInputIndex = 0;
+ intermediateEdges.push_back(rotatedDataToMulEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC sinToMulEdge = {};
+ sinToMulEdge.FromNodeIndex = gatherSinOpIndex;
+ sinToMulEdge.FromNodeOutputIndex = 0;
+ sinToMulEdge.ToNodeIndex = mulSinOpIndex;
+ sinToMulEdge.ToNodeInputIndex = 1;
+ intermediateEdges.push_back(sinToMulEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC rotatedSinToMulEdge = {};
+ rotatedSinToMulEdge.FromNodeIndex = mulSinOpIndex;
+ rotatedSinToMulEdge.FromNodeOutputIndex = 0;
+ rotatedSinToMulEdge.ToNodeIndex = mulSignOpIndex;
+ rotatedSinToMulEdge.ToNodeInputIndex = 0;
+ intermediateEdges.push_back(rotatedSinToMulEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC signToMulEdge = {};
+ signToMulEdge.FromNodeIndex = signRangeOpIndex;
+ signToMulEdge.FromNodeOutputIndex = 0;
+ signToMulEdge.ToNodeIndex = mulSignOpIndex;
+ signToMulEdge.ToNodeInputIndex = 1;
+ intermediateEdges.push_back(signToMulEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC nonRotatedCosToAddEdge = {};
+ nonRotatedCosToAddEdge.FromNodeIndex = mulCosOpIndex;
+ nonRotatedCosToAddEdge.FromNodeOutputIndex = 0;
+ nonRotatedCosToAddEdge.ToNodeIndex = addOpIndex;
+ nonRotatedCosToAddEdge.ToNodeInputIndex = 0;
+ intermediateEdges.push_back(nonRotatedCosToAddEdge);
+
+ DML_INTERMEDIATE_GRAPH_EDGE_DESC rotatedSinToAddEdge = {};
+ rotatedSinToAddEdge.FromNodeIndex = mulSignOpIndex;
+ rotatedSinToAddEdge.FromNodeOutputIndex = 0;
+ rotatedSinToAddEdge.ToNodeIndex = addOpIndex;
+ rotatedSinToAddEdge.ToNodeInputIndex = 1;
+ intermediateEdges.push_back(rotatedSinToAddEdge);
+
+ DML_OUTPUT_GRAPH_EDGE_DESC addToOutputEdge = {};
+ addToOutputEdge.FromNodeIndex = addOpIndex;
+ addToOutputEdge.FromNodeOutputIndex = 0;
+ addToOutputEdge.GraphOutputIndex = 0;
+ outputEdges.push_back(addToOutputEdge);
+
+ MLOperatorGraphDesc operatorGraphDesc = {};
+ operatorGraphDesc.inputEdgeCount = gsl::narrow_cast(inputEdges.size());
+ operatorGraphDesc.inputEdges = inputEdges.data();
+ operatorGraphDesc.intermediateEdgeCount = gsl::narrow_cast(intermediateEdges.size());
+ operatorGraphDesc.intermediateEdges = intermediateEdges.data();
+ operatorGraphDesc.outputEdgeCount = gsl::narrow_cast(outputEdges.size());
+ operatorGraphDesc.outputEdges = outputEdges.data();
+ operatorGraphDesc.nodeCount = gsl::narrow_cast(opDescs.size());
+ operatorGraphDesc.nodesAsOpDesc = opDescs.data();
+
+ SetDmlOperatorGraphDesc(std::move(operatorGraphDesc), kernelInfo);
+ }
+};
+
+DML_OP_DEFINE_CREATION_FUNCTION(RotaryEmbedding, DmlOperatorRotaryEmbedding);
+
+} // 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 30bc6e5e27..28360f09bc 100644
--- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp
+++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp
@@ -510,6 +510,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(BitwiseAnd);
DML_OP_EXTERN_CREATION_FUNCTION(BitwiseOr);
DML_OP_EXTERN_CREATION_FUNCTION(BitwiseXor);
DML_OP_EXTERN_CREATION_FUNCTION(BitwiseNot);
+DML_OP_EXTERN_CREATION_FUNCTION(RotaryEmbedding);
DML_OP_EXTERN_QUERY_FUNCTION(MaxPool);
DML_OP_EXTERN_QUERY_FUNCTION(Slice);
@@ -527,6 +528,7 @@ DML_OP_EXTERN_QUERY_FUNCTION(Attention);
constexpr static std::array typeNameListDefault = {"T"};
constexpr static std::array typeNameListDefaultV = {"V"};
constexpr static std::array typeNameListAttention = {"T", "M"};
+constexpr static std::array typeNameListRotaryEmbedding = {"T", "M"};
constexpr static std::array typeNameListTwo = { "T1", "T2" };
constexpr static std::array typeNameListLayerNorm = { "T", "U" };
constexpr static std::array typeNameListLayerNormContrib = { "T", "V" };
@@ -597,6 +599,7 @@ constexpr static std::array supportedTypeListShape
constexpr static std::array supportedTypeListSize = {SupportedTensorDataTypes::All, SupportedTensorDataTypes::Int64};
constexpr static std::array supportedTypeListQLinearSigmoid = {SupportedTensorDataTypes::UInt8 | SupportedTensorDataTypes::Int8};
constexpr static std::array supportedTypeListAttention = {SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::Int32};
+constexpr static std::array supportedTypeListRotaryEmbedding = {SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::Int64};
constexpr static std::array supportedTypeListGroupNorm = {SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::Float16to32};
constexpr static std::array supportedTypeListNonZero = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Ints8Bit | SupportedTensorDataTypes::Ints16Bit | SupportedTensorDataTypes::Ints32Bit | SupportedTensorDataTypes::Bool};
@@ -1006,6 +1009,7 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation
{REG_INFO_MS( 1, QLinearSigmoid, typeNameListDefault, supportedTypeListQLinearSigmoid, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryQLinearSigmoid)},
{REG_INFO_MS( 1, Attention, typeNameListAttention, supportedTypeListAttention, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryAttention)},
{REG_INFO_MS( 1, MultiHeadAttention, typeNameListAttention, supportedTypeListAttention, DmlGraphSupport::Supported)},
+ {REG_INFO_MS( 1, RotaryEmbedding, typeNameListRotaryEmbedding, supportedTypeListRotaryEmbedding, DmlGraphSupport::Supported)},
{REG_INFO( 10, IsInf, typeNameListTwo, supportedTypeListIsInf, DmlGraphSupport::Supported)},
{REG_INFO( 10, Mod, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported)},
diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h
index dac128f92a..e9591cfce6 100644
--- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h
+++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h
@@ -122,6 +122,7 @@ namespace AttrName
static constexpr const char* GraphFusedActivation = "activation";
static constexpr const char* GraphFusedAxis = "activation_axis";
+ static constexpr const char* Interleaved = "interleaved";
} // namespace AttrName
diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h
index 485e20c1df..f7e545d9d9 100644
--- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h
+++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h
@@ -1584,6 +1584,7 @@ using ShapeInferenceHelper_DequantizeLinear = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_QLinearSigmoid = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_Attention = AttentionHelper;
using ShapeInferenceHelper_MultiHeadAttention = MultiHeadAttentionHelper;
+using ShapeInferenceHelper_RotaryEmbedding = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_Sign = GetBroadcastedOutputShapeHelper;
using ShapeInferenceHelper_IsNaN = GetBroadcastedOutputShapeHelper;
using ShapeInferenceHelper_Erf = GetBroadcastedOutputShapeHelper;
diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h
index c1e525400b..e18ba31def 100644
--- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h
+++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h
@@ -437,6 +437,7 @@ namespace OperatorHelper
static const int sc_sinceVer_BiasAdd = 1;
static const int sc_sinceVer_QuickGelu = 1;
static const int sc_sinceVer_GroupNorm = 1;
+ static const int sc_sinceVer_RotaryEmbedding = 1;
} // namespace MsftOperatorSet1
} // namespace OperatorHelper
diff --git a/onnxruntime/python/onnxruntime_pybind_iobinding.cc b/onnxruntime/python/onnxruntime_pybind_iobinding.cc
index 7638a12bb8..59d5a77bfb 100644
--- a/onnxruntime/python/onnxruntime_pybind_iobinding.cc
+++ b/onnxruntime/python/onnxruntime_pybind_iobinding.cc
@@ -60,8 +60,6 @@ void addIoBindingMethods(pybind11::module& m) {
})
// This binds input as a Tensor that wraps memory pointer along with the OrtMemoryInfo
.def("bind_input", [](SessionIOBinding* io_binding, const std::string& name, const OrtDevice& device, py::object& element_type, const std::vector& shape, int64_t data_ptr) -> void {
- ORT_ENFORCE(data_ptr != 0, "Pointer to data memory is not valid");
-
PyArray_Descr* dtype;
if (!PyArray_DescrConverter(element_type.ptr(), &dtype)) {
throw std::runtime_error("Not a valid numpy type");
diff --git a/onnxruntime/python/tools/symbolic_shape_infer.py b/onnxruntime/python/tools/symbolic_shape_infer.py
index 5899c4fcfc..69f8530dff 100755
--- a/onnxruntime/python/tools/symbolic_shape_infer.py
+++ b/onnxruntime/python/tools/symbolic_shape_infer.py
@@ -154,6 +154,8 @@ class SymbolicShapeInference:
"MatMulInteger16": self._infer_MatMulInteger,
"MaxPool": self._infer_Pool,
"Max": self._infer_symbolic_compute_ops,
+ "MemcpyFromHost": self._pass_on_shape_and_type,
+ "MemcpyToHost": self._pass_on_shape_and_type,
"Min": self._infer_symbolic_compute_ops,
"Mul": self._infer_symbolic_compute_ops,
"NonMaxSuppression": self._infer_NonMaxSuppression,
diff --git a/onnxruntime/python/tools/transformers/float16.py b/onnxruntime/python/tools/transformers/float16.py
index 222f5f5e27..95e7437493 100644
--- a/onnxruntime/python/tools/transformers/float16.py
+++ b/onnxruntime/python/tools/transformers/float16.py
@@ -145,7 +145,8 @@ DEFAULT_OP_BLOCK_LIST = [
# Some operators has data type fixed as float for some inputs. Key is op_type, value is list of input indices
-ALWAYS_FLOAT_INPUTS = {"Resize": [2], "GroupNorm": [1, 2]}
+# Note that DirectML allows float16 gamma and beta in GroupNorm. Use force_fp16_inputs parameter could overwrite this.
+ALWAYS_FLOAT_INPUTS = {"Resize": [2], "GroupNorm": [1, 2], "SkipGroupNorm": [1, 2]}
class InitializerTracker:
diff --git a/onnxruntime/python/tools/transformers/fusion_group_norm.py b/onnxruntime/python/tools/transformers/fusion_group_norm.py
index cd7dc7017c..c718d2c27e 100644
--- a/onnxruntime/python/tools/transformers/fusion_group_norm.py
+++ b/onnxruntime/python/tools/transformers/fusion_group_norm.py
@@ -82,23 +82,11 @@ class FusionGroupNorm(Fusion):
return
instance_norm_scale = self.model.get_constant_value(instance_norm.input[1])
- if instance_norm_scale is None:
+ if instance_norm_scale is None or len(instance_norm_scale.shape) != 1:
return
+
instance_norm_bias = self.model.get_constant_value(instance_norm.input[2])
- if instance_norm_bias is None:
- return
-
- # Only groups=32 is supported in GroupNorm kernel. Check the scale and bias is 1D tensor with shape [32].
- if not (len(instance_norm_scale.shape) == 1 and instance_norm_scale.shape[0] == 32):
- logger.debug(
- "Skip GroupNorm fusion since scale shape is expected to be [32], Got %s", str(instance_norm_scale.shape)
- )
- return
-
- if not (len(instance_norm_bias.shape) == 1 and instance_norm_bias.shape[0] == 32):
- logger.debug(
- "Skip GroupNorm fusion since bias shape is expected to be [32], Got %s", str(instance_norm_bias.shape)
- )
+ if instance_norm_bias is None or instance_norm_scale.shape != instance_norm_scale.shape:
return
if not np.allclose(np.ones_like(instance_norm_scale), instance_norm_scale):
@@ -108,10 +96,6 @@ class FusionGroupNorm(Fusion):
group_norm_name = self.model.create_node_name("GroupNorm", name_prefix="GroupNorm")
- if weight_elements not in [320, 640, 960, 1280, 1920, 2560, 128, 256, 512]:
- logger.info("Skip GroupNorm fusion since channels=%d is not supported.", weight_elements)
- return
-
self.add_initializer(
name=group_norm_name + "_gamma",
data_type=TensorProto.FLOAT,
diff --git a/onnxruntime/python/tools/transformers/fusion_options.py b/onnxruntime/python/tools/transformers/fusion_options.py
index 8c80fcad0a..b9b92d2fe8 100644
--- a/onnxruntime/python/tools/transformers/fusion_options.py
+++ b/onnxruntime/python/tools/transformers/fusion_options.py
@@ -61,6 +61,7 @@ class FusionOptions:
if model_type in ["unet", "vae", "clip"]:
self.enable_nhwc_conv = True
self.enable_group_norm = True
+ self.enable_skip_group_norm = True
self.enable_bias_splitgelu = True
self.enable_packed_qkv = True
self.enable_packed_kv = True
@@ -116,6 +117,8 @@ class FusionOptions:
options.enable_nhwc_conv = False
if args.disable_group_norm:
options.enable_group_norm = False
+ if args.disable_skip_group_norm:
+ options.enable_skip_group_norm = False
if args.disable_bias_splitgelu:
options.enable_bias_splitgelu = False
if args.disable_packed_qkv:
@@ -250,6 +253,14 @@ class FusionOptions:
)
parser.set_defaults(disable_group_norm=False)
+ parser.add_argument(
+ "--disable_skip_group_norm",
+ required=False,
+ action="store_true",
+ help="not fuse Add + GroupNorm to SkipGroupNorm. Only works for model_type=unet or vae",
+ )
+ parser.set_defaults(disable_skip_group_norm=False)
+
parser.add_argument(
"--disable_packed_kv",
required=False,
diff --git a/onnxruntime/python/tools/transformers/fusion_rotary_attention.py b/onnxruntime/python/tools/transformers/fusion_rotary_attention.py
index ceee836e33..de89b35366 100644
--- a/onnxruntime/python/tools/transformers/fusion_rotary_attention.py
+++ b/onnxruntime/python/tools/transformers/fusion_rotary_attention.py
@@ -29,7 +29,13 @@ class FusionRotaryAttention(FusionAttention):
hidden_size,
num_heads,
use_multi_head_attention=True,
- search_op_types=["SimplifiedLayerNormalization", "SkipSimplifiedLayerNormalization", "Add"],
+ search_op_types=[
+ "SimplifiedLayerNormalization",
+ "SkipSimplifiedLayerNormalization",
+ "LayerNormalization",
+ "SkipLayerNormalization",
+ "Add",
+ ],
)
def create_mha_node(
@@ -318,7 +324,7 @@ class FusionRotaryAttention(FusionAttention):
return True
def fuse(self, normalize_node, input_name_to_nodes, output_name_to_node):
- if normalize_node.op_type != "SkipSimplifiedLayerNormalization" and normalize_node.op_type != "Add":
+ if normalize_node.op_type not in {"SkipSimplifiedLayerNormalization", "SkipLayerNormalization", "Add"}:
return
# qkv_nodes_1 is for LLaMA-2 Microsoft
diff --git a/onnxruntime/python/tools/transformers/fusion_skip_group_norm.py b/onnxruntime/python/tools/transformers/fusion_skip_group_norm.py
new file mode 100644
index 0000000000..df80acbd97
--- /dev/null
+++ b/onnxruntime/python/tools/transformers/fusion_skip_group_norm.py
@@ -0,0 +1,255 @@
+# -------------------------------------------------------------------------
+# Copyright (c) Microsoft Corporation. All rights reserved.
+# Licensed under the MIT License.
+# --------------------------------------------------------------------------
+from logging import getLogger
+from typing import List
+
+from fusion_base import Fusion
+from fusion_utils import NumpyHelper
+from onnx import helper
+from onnx_model import OnnxModel
+
+logger = getLogger(__name__)
+
+
+class FusionSkipGroupNorm(Fusion):
+ """
+ Fuse Add + GroupNorm into one node: SkipGroupNorm.
+ """
+
+ def __init__(self, model: OnnxModel):
+ super().__init__(model, "SkipGroupNorm", "GroupNorm")
+ # Update shape inference is needed since other fusions might add new edge which does not have shape info yet.
+ self.shape_infer_helper = self.model.infer_runtime_shape(update=True)
+
+ if self.shape_infer_helper is None:
+ logger.warning("SkipGroupNorm fusion will be skipped since symbolic shape inference disabled or failed.")
+
+ def create_transpose_node(self, input_name: str, perm: List[int], output_name=None):
+ """Append a Transpose node after an input"""
+ node_name = self.model.create_node_name("Transpose")
+ if output_name is None:
+ output_name = node_name + "_out" + "-" + input_name
+ transpose_node = helper.make_node("Transpose", inputs=[input_name], outputs=[output_name], name=node_name)
+ transpose_node.attribute.extend([helper.make_attribute("perm", perm)])
+ return transpose_node
+
+ def get_skip_index(self, add, is_channel_last: bool):
+ """Add has two inputs. This classifies which input is skip based on shape info (skip allows broadcast)."""
+ skip = -1
+ broadcast = False
+
+ assert self.shape_infer_helper is not None
+ shape_a = self.shape_infer_helper.get_edge_shape(add.input[0])
+ shape_b = self.shape_infer_helper.get_edge_shape(add.input[1])
+ assert shape_a is not None and shape_b is not None
+
+ if len(shape_a) == 4 and len(shape_b) == 4:
+ if shape_a == shape_b:
+ skip = 1
+ else:
+ c = 3 if is_channel_last else 1
+ h = 1 if is_channel_last else 2
+ w = 2 if is_channel_last else 3
+ if shape_a[0] == shape_b[0] and shape_a[c] == shape_b[c]:
+ if shape_b[h] == 1 and shape_b[w] == 1:
+ skip = 1
+ broadcast = True
+ elif shape_a[h] == 1 and shape_a[w] == 1:
+ skip = 0
+ broadcast = True
+
+ if skip < 0:
+ logger.debug(
+ "skip SkipGroupNorm fusion since shape of Add inputs (%s, %s) are not expected",
+ add.input[0],
+ add.input[1],
+ )
+ return skip, broadcast
+
+ def has_multiple_consumers(self, output_name, input_name_to_nodes):
+ """Whether an output has multiple consumers (like graph output or more than one children nodes)"""
+ return self.model.find_graph_output(output_name) is not None or (
+ output_name in input_name_to_nodes and len(input_name_to_nodes[output_name]) > 1
+ )
+
+ def remove_if_safe(self, node, input_name_to_nodes):
+ """Remove a node if it is safe (only one children, and not graph output)"""
+ if not self.has_multiple_consumers(node.output[0], input_name_to_nodes):
+ self.nodes_to_remove.extend([node])
+
+ def is_bias_1d(self, bias_name: str):
+ """Whether bias is an initializer of one dimension"""
+ initializer = self.model.get_initializer(bias_name)
+ if initializer is None:
+ return False
+
+ bias_weight = NumpyHelper.to_array(initializer)
+ if bias_weight is None:
+ logger.debug("Bias weight not found")
+ return False
+
+ if len(bias_weight.shape) != 1:
+ logger.debug("Bias weight is not 1D")
+ return False
+ return True
+
+ def match_bias_path(self, node, input_name_to_nodes, output_name_to_node):
+ """
+ Match the bias graph pattern from an Transpose node after Reshape node like in below example.
+ It checks whether the bias is 1D initializer. If so, remove Add and redirect MatMul output to Reshape.
+ """
+ # Before Fusion:
+ # MatMul (bias)
+ # \ / (shape)
+ # Add /
+ # \ /
+ # (a) Reshape
+ # \ |
+ # Transpose([0, 3, 1, 2]) Transpose([0, 3, 1, 2]) --- the start node, this func only handles the above nodes.
+ # \ /
+ # Add
+ # / \
+ # (c) Transpose([0,2,3,1])
+ # |
+ # GroupNorm
+ # |
+ # (d)
+ #
+ # After Fusion (the nodes below Reshape is handled in the fuse function):
+ # MatMul (shape)
+ # \ /
+ # (a) Reshape
+ # \ /
+ # SkipGroupNorm
+ # / \
+ # (d) Transpose([0, 3, 1, 2])
+ # \
+ # (c)
+
+ add_input_index = []
+ bias_nodes = self.model.match_parent_path(
+ node, ["Reshape", "Add", "MatMul"], [0, 0, None], output_name_to_node, add_input_index
+ )
+ if bias_nodes is None:
+ return None
+
+ (reshape, add_bias, matmul) = bias_nodes
+ bias = bias_nodes[1].input[1 - add_input_index[0]]
+ if not self.is_bias_1d(bias):
+ return None
+
+ reshape.input[0] = matmul.output[0]
+ self.remove_if_safe(add_bias, input_name_to_nodes)
+
+ return bias
+
+ def match_transpose_from_nhwc(self, output_name, input_name_to_nodes, output_name_to_node):
+ """Match whether an output is from a Transpose(perm=[0,3,1,2]) node."""
+ parent = output_name_to_node[output_name] if output_name in output_name_to_node else None
+ if parent is not None and parent.op_type == "Transpose":
+ permutation = OnnxModel.get_node_attribute(parent, "perm")
+ if permutation == [0, 3, 1, 2]:
+ self.remove_if_safe(parent, input_name_to_nodes)
+ return parent
+ return None
+
+ def fuse(self, node, input_name_to_nodes, output_name_to_node):
+ # This fusion requires shape information, so skip it if shape is not available.
+ if self.shape_infer_helper is None:
+ return
+
+ # Before Fusion:
+ # (a) (b)
+ # \ /
+ # Add
+ # /\
+ # (c) Transpose([0,2,3,1])
+ # \
+ # GroupNorm
+ # |
+ # (d)
+ #
+ # After Fusion:
+ # (a) (b)
+ # \ /
+ # Transpose([0,2,3,1]) Transpose([0,2,3,1])
+ # \ /
+ # SkipGroupNorm
+ # / \
+ # / Transpose([0, 3, 1, 2])
+ # / \
+ # (d) (c)
+ nodes = self.model.match_parent_path(node, ["Transpose", "Add"], [0, 0], output_name_to_node)
+ if nodes is None:
+ return
+
+ (transpose, add) = nodes
+ if transpose in self.nodes_to_remove or add in self.nodes_to_remove:
+ return
+
+ if self.has_multiple_consumers(transpose.output[0], input_name_to_nodes):
+ return
+
+ permutation = OnnxModel.get_node_attribute(transpose, "perm")
+ if permutation != [0, 2, 3, 1]:
+ return
+
+ inputs = []
+ bias = None
+ for i in range(2):
+ matched_transpose = self.match_transpose_from_nhwc(add.input[i], input_name_to_nodes, output_name_to_node)
+ if matched_transpose:
+ # When there is an Transpose node before Add (see examples in match_bias_path), we do not need to
+ # insert another Transpose node. The existing Transpose node will be removed in prune_graph if it
+ # has only one consumer.
+ inputs.append(matched_transpose.input[0])
+ # See whether it match bias pattern.
+ if bias is None:
+ bias = self.match_bias_path(matched_transpose, input_name_to_nodes, output_name_to_node)
+ else:
+ # Otherwise, insert a Transpose node before Add.
+ new_transpose = self.create_transpose_node(add.input[i], [0, 2, 3, 1])
+ self.model.add_node(new_transpose, self.this_graph_name)
+ inputs.append(new_transpose.output[0])
+
+ skip, broadcast = self.get_skip_index(add, is_channel_last=False)
+ if skip < 0:
+ return
+
+ inputs = [inputs[1 - skip], node.input[1], node.input[2], inputs[skip]]
+ if bias:
+ inputs = [*inputs, bias]
+
+ outputs = node.output
+
+ new_node_name = self.model.create_node_name(self.fused_op_type, name_prefix="SkipGroupNorm")
+ if self.has_multiple_consumers(add.output[0], input_name_to_nodes):
+ add_out_name = new_node_name + "_add_out"
+ outputs.append(add_out_name)
+
+ # Insert a Transpose node after add output.
+ add_out_transpose = self.create_transpose_node(add_out_name, [0, 3, 1, 2], add.output[0])
+ self.model.add_node(add_out_transpose, self.this_graph_name)
+
+ skip_group_norm = helper.make_node(
+ self.fused_op_type,
+ inputs=inputs,
+ outputs=outputs,
+ name=new_node_name,
+ )
+ skip_group_norm.domain = "com.microsoft"
+
+ self.increase_counter(
+ f"SkipGroupNorm(add_out={int(len(outputs) > 1)} bias={int(bias is not None)} broadcast={int(broadcast)})"
+ )
+
+ # Pass attributes from GroupNorm node to SkipGroupNorm
+ for att in node.attribute:
+ skip_group_norm.attribute.extend([att])
+
+ self.nodes_to_remove.extend([add, transpose, node])
+ self.nodes_to_add.append(skip_group_norm)
+ self.node_name_to_graph_name[skip_group_norm.name] = self.this_graph_name
+ self.prune_graph = True
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_txt2img.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_txt2img.py
index d6de5c45a5..fb051ac1ed 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_txt2img.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_txt2img.py
@@ -61,6 +61,9 @@ if __name__ == "__main__":
_, shared_device_memory = cudart.cudaMalloc(max_device_memory)
pipeline.backend.activate_engines(shared_device_memory)
+ if engine_type == EngineType.ORT_CUDA and args.enable_vae_slicing:
+ pipeline.backend.enable_vae_slicing()
+
pipeline.load_resources(image_height, image_width, batch_size)
def run_inference(warmup=False):
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_txt2img_xl.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_txt2img_xl.py
index efc87a207d..16e776a082 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_txt2img_xl.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_txt2img_xl.py
@@ -70,6 +70,14 @@ def run_demo():
base.backend.activate_engines(shared_device_memory)
refiner.backend.activate_engines(shared_device_memory)
+ if engine_type == EngineType.ORT_CUDA:
+ enable_vae_slicing = args.enable_vae_slicing
+ if batch_size > 4 and not enable_vae_slicing:
+ print("Updating enable_vae_slicing to be True to avoid cuDNN error for batch size > 4.")
+ enable_vae_slicing = True
+ if enable_vae_slicing:
+ refiner.backend.enable_vae_slicing()
+
base.load_resources(image_height, image_width, batch_size)
refiner.load_resources(image_height, image_width, batch_size)
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_utils.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_utils.py
index 3996c8c325..e65efd2c53 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_utils.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/demo_utils.py
@@ -68,7 +68,7 @@ def parse_arguments(is_xl: bool, description: str):
"--scheduler",
type=str,
default="DDIM",
- choices=["DDIM", "EulerA", "UniPC"],
+ choices=["DDIM", "UniPC"] if is_xl else ["DDIM", "EulerA", "UniPC"],
help="Scheduler for diffusion process",
)
@@ -145,6 +145,9 @@ def parse_arguments(is_xl: bool, description: str):
parser.add_argument("--seed", type=int, default=None, help="Seed for random generator to get consistent results.")
parser.add_argument("--disable-cuda-graph", action="store_true", help="Disable cuda graph.")
+ group = parser.add_argument_group("Options for ORT_CUDA engine only")
+ group.add_argument("--enable-vae-slicing", action="store_true", help="True will feed only one image to VAE once.")
+
# TensorRT only options
group = parser.add_argument_group("Options for TensorRT (--engine=TRT) only")
group.add_argument("--onnx-refit-dir", help="ONNX models to load the weights from.")
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/diffusion_models.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/diffusion_models.py
index dc777e2693..4a2e9eb344 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/diffusion_models.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/diffusion_models.py
@@ -303,7 +303,15 @@ class BaseModel:
"""
return []
- def optimize_ort(self, input_onnx_path, optimized_onnx_path, to_fp16=True, fp32_op_list=None, optimize_by_ort=True):
+ def optimize_ort(
+ self,
+ input_onnx_path,
+ optimized_onnx_path,
+ to_fp16=True,
+ fp32_op_list=None,
+ optimize_by_ort=True,
+ optimize_by_fusion=True,
+ ):
optimizer = self.get_ort_optimizer()
optimizer.optimize(
input_onnx_path,
@@ -312,6 +320,7 @@ class BaseModel:
keep_io_types=self.fp32_input_output_names(),
fp32_op_list=fp32_op_list,
optimize_by_ort=optimize_by_ort,
+ optimize_by_fusion=optimize_by_fusion,
)
def optimize_trt(self, input_onnx_path, optimized_onnx_path):
@@ -471,7 +480,15 @@ class CLIP(BaseModel):
onnx_model.add_node(cast_node)
onnx_model.save_model_to_file(optimized_onnx_path, use_external_data_format=use_external_data_format)
- def optimize_ort(self, input_onnx_path, optimized_onnx_path, to_fp16=True, fp32_op_list=None, optimize_by_ort=True):
+ def optimize_ort(
+ self,
+ input_onnx_path,
+ optimized_onnx_path,
+ to_fp16=True,
+ fp32_op_list=None,
+ optimize_by_ort=True,
+ optimize_by_fusion=True,
+ ):
optimizer = self.get_ort_optimizer()
if not self.output_hidden_state:
@@ -483,8 +500,9 @@ class CLIP(BaseModel):
fp32_op_list=fp32_op_list,
keep_outputs=["text_embeddings"],
optimize_by_ort=optimize_by_ort,
+ optimize_by_fusion=optimize_by_fusion,
)
- else:
+ elif optimize_by_fusion:
with tempfile.TemporaryDirectory() as tmp_dir:
# Save to a temporary file so that we can load it with Onnx Runtime.
logger.info("Saving a temporary model to add hidden_states to graph output ...")
@@ -500,7 +518,19 @@ class CLIP(BaseModel):
fp32_op_list=fp32_op_list,
keep_outputs=["text_embeddings", "hidden_states"],
optimize_by_ort=optimize_by_ort,
+ optimize_by_fusion=optimize_by_fusion,
)
+ else: # input is optimized model, there is no need to add hidden states.
+ optimizer.optimize(
+ input_onnx_path,
+ optimized_onnx_path,
+ float16=to_fp16,
+ keep_io_types=[],
+ fp32_op_list=fp32_op_list,
+ keep_outputs=["text_embeddings", "hidden_states"],
+ optimize_by_ort=optimize_by_ort,
+ optimize_by_fusion=optimize_by_fusion,
+ )
def optimize_trt(self, input_onnx_path, optimized_onnx_path):
onnx_graph = onnx.load(input_onnx_path)
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/diffusion_schedulers.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/diffusion_schedulers.py
index 13c450a517..ec3041e134 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/diffusion_schedulers.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/diffusion_schedulers.py
@@ -695,6 +695,7 @@ class UniPCMultistepScheduler:
self,
original_samples: torch.FloatTensor,
noise: torch.FloatTensor,
+ idx,
timesteps: torch.IntTensor,
) -> torch.FloatTensor:
# Make sure alphas_cumprod and timestep have same device and dtype as original_samples
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/engine_builder.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/engine_builder.py
index 029125c639..dfdfa007d7 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/engine_builder.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/engine_builder.py
@@ -73,6 +73,10 @@ class EngineBuilder:
self.models = {}
self.engines = {}
self.torch_models = {}
+ self.use_vae_slicing = False
+
+ def enable_vae_slicing(self):
+ self.use_vae_slicing = True
def teardown(self):
for engine in self.engines.values():
@@ -84,9 +88,9 @@ class EngineBuilder:
model_name += "_inpaint"
return model_name
- def get_onnx_path(self, model_name, onnx_dir, opt=True):
+ def get_onnx_path(self, model_name, onnx_dir, opt=True, suffix=""):
engine_name = self.engine_type.name.lower()
- directory_name = self.get_cached_model_name(model_name) + (f".{engine_name}" if opt else "")
+ directory_name = self.get_cached_model_name(model_name) + (f".{engine_name}" if opt else "") + suffix
onnx_model_dir = os.path.join(onnx_dir, directory_name)
os.makedirs(onnx_model_dir, exist_ok=True)
return os.path.join(onnx_model_dir, "model.onnx")
@@ -160,11 +164,12 @@ class EngineBuilder:
for model_name, obj in self.models.items():
if model_name == "vae" and self.vae_torch_fallback:
continue
+ slice_size = 1 if (model_name == "vae" and self.use_vae_slicing) else batch_size
self.engines[model_name].allocate_buffers(
- shape_dict=obj.get_shape_dict(batch_size, image_height, image_width), device=self.torch_device
+ shape_dict=obj.get_shape_dict(slice_size, image_height, image_width), device=self.torch_device
)
- def vae_decode(self, latents):
+ def _vae_decode(self, latents):
if self.vae_torch_fallback:
if not self.custom_fp16_vae:
latents = latents.to(dtype=torch.float32)
@@ -175,6 +180,14 @@ class EngineBuilder:
return images
+ def vae_decode(self, latents):
+ if self.use_vae_slicing:
+ # The output tensor points to same buffer. Need clone it to avoid overwritten.
+ decoded_slices = [self._vae_decode(z_slice).clone() for z_slice in latents.split(1)]
+ return torch.cat(decoded_slices)
+
+ return self._vae_decode(latents)
+
def get_engine_paths(work_dir: str, pipeline_info: PipelineInfo, engine_type: EngineType):
root_dir = work_dir or "."
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/engine_builder_ort_cuda.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/engine_builder_ort_cuda.py
index 11a39b0dec..07c675b2ed 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/engine_builder_ort_cuda.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/engine_builder_ort_cuda.py
@@ -144,7 +144,7 @@ class OrtCudaEngineBuilder(EngineBuilder):
self._configure(
"unetxl",
onnx_opset_version=onnx_opset_version,
- use_cuda_graph=False, # TODO: fix Runtime Error with cuda graph
+ use_cuda_graph=self.use_cuda_graph,
)
self._configure(
@@ -164,6 +164,7 @@ class OrtCudaEngineBuilder(EngineBuilder):
opt_batch_size: int = 1,
force_engine_rebuild: bool = False,
device_id: int = 0,
+ save_fp32_intermediate_model=False,
):
self.torch_device = torch.device("cuda", device_id)
self.load_models(framework_model_dir)
@@ -195,7 +196,9 @@ class OrtCudaEngineBuilder(EngineBuilder):
continue
onnx_path = self.get_onnx_path(model_name, onnx_dir, opt=False)
- onnx_opt_path = self.get_onnx_path(model_name, engine_dir, opt=True)
+ onnx_fp32_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp32")
+ onnx_fp16_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp16")
+ onnx_opt_path = onnx_fp16_path if self.model_config[model_name].fp16 else onnx_fp32_path
if not os.path.exists(onnx_opt_path):
if not os.path.exists(onnx_path):
print("----")
@@ -225,17 +228,41 @@ class OrtCudaEngineBuilder(EngineBuilder):
else:
logger.info("Found cached model: %s", onnx_path)
- # Run graph optimization and convert to mixed precision (computation in FP16)
+ # Generate fp32 optimized model.
+ # If final target is fp16 model, we save fp32 optimized model so that it is easy to tune
+ # fp16 conversion. That could save a lot of time in developing.
+ use_fp32_intermediate = save_fp32_intermediate_model and self.model_config[model_name].fp16
+ if use_fp32_intermediate:
+ if not os.path.exists(onnx_fp32_path):
+ print("------")
+ logger.info("Generating optimized model: %s", onnx_fp32_path)
+
+ # There is risk that some ORT fused ops fp32 only. So far, we have not encountered such issue.
+ model_obj.optimize_ort(
+ onnx_path,
+ onnx_fp32_path,
+ to_fp16=False,
+ fp32_op_list=self.model_config[model_name].force_fp32_ops,
+ optimize_by_ort=self.model_config[model_name].optimize_by_ort,
+ )
+ else:
+ logger.info("Found cached optimized model: %s", onnx_fp32_path)
+
+ # Generate the final optimized model.
if not os.path.exists(onnx_opt_path):
print("------")
logger.info("Generating optimized model: %s", onnx_opt_path)
+ # When there is fp32 intermediate optimized model, this will just convert model from fp32 to fp16.
+ optimize_by_ort = False if use_fp32_intermediate else self.model_config[model_name].optimize_by_ort
+
model_obj.optimize_ort(
- onnx_path,
+ onnx_fp32_path if use_fp32_intermediate else onnx_path,
onnx_opt_path,
to_fp16=self.model_config[model_name].fp16,
fp32_op_list=self.model_config[model_name].force_fp32_ops,
- optimize_by_ort=self.model_config[model_name].optimize_by_ort,
+ optimize_by_ort=optimize_by_ort,
+ optimize_by_fusion=not use_fp32_intermediate,
)
else:
logger.info("Found cached optimized model: %s", onnx_opt_path)
@@ -245,7 +272,9 @@ class OrtCudaEngineBuilder(EngineBuilder):
if model_name == "vae" and self.vae_torch_fallback:
continue
- onnx_opt_path = self.get_onnx_path(model_name, engine_dir, opt=True)
+ onnx_fp32_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp32")
+ onnx_fp16_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp16")
+ onnx_opt_path = onnx_fp16_path if self.model_config[model_name].fp16 else onnx_fp32_path
use_cuda_graph = self.model_config[model_name].use_cuda_graph
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/ort_optimizer.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/ort_optimizer.py
index 2078f8d1a4..4b48396b6c 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/ort_optimizer.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/ort_optimizer.py
@@ -60,34 +60,37 @@ class OrtStableDiffusionOptimizer:
fp32_op_list=None,
keep_outputs=None,
optimize_by_ort=True,
+ optimize_by_fusion=True,
+ final_target_float16=True,
):
"""Optimize onnx model using ONNX Runtime transformers optimizer"""
logger.info(f"Optimize {input_fp32_onnx_path}...")
- fusion_options = FusionOptions(self.model_type)
- if self.model_type in ["unet"] and not float16:
- fusion_options.enable_packed_kv = False
- fusion_options.enable_packed_qkv = False
- m = optimize_model(
- input_fp32_onnx_path,
- model_type=self.model_type,
- num_heads=0, # will be deduced from graph
- hidden_size=0, # will be deduced from graph
- opt_level=0,
- optimization_options=fusion_options,
- use_gpu=True,
- )
+ if optimize_by_fusion:
+ fusion_options = FusionOptions(self.model_type)
+
+ # It is allowed float16=False and final_target_float16=True, for using fp32 as intermediate optimization step.
+ # For rare fp32 use case, we can disable packed kv/qkv since there is no fp32 TRT fused attention kernel.
+ if self.model_type in ["unet"] and not final_target_float16:
+ fusion_options.enable_packed_kv = False
+ fusion_options.enable_packed_qkv = False
+
+ m = optimize_model(
+ input_fp32_onnx_path,
+ model_type=self.model_type,
+ num_heads=0, # will be deduced from graph
+ hidden_size=0, # will be deduced from graph
+ opt_level=0,
+ optimization_options=fusion_options,
+ use_gpu=True,
+ )
+ else:
+ model = onnx.load_model(input_fp32_onnx_path, load_external_data=True)
+ m = self.model_type_class_mapping[self.model_type](model)
if keep_outputs:
m.prune_graph(outputs=keep_outputs)
- if float16:
- logger.info("Convert to float16 ...")
- m.convert_float_to_float16(
- keep_io_types=keep_io_types,
- op_block_list=fp32_op_list,
- )
-
use_external_data_format = m.model.ByteSize() >= onnx.checker.MAXIMUM_PROTOBUF
# Note that ORT < 1.16 could not save model larger than 2GB.
@@ -100,6 +103,13 @@ class OrtStableDiffusionOptimizer:
if optimize_by_ort and (version.parse(ort_version) >= version.parse("1.16.0") or not use_external_data_format):
m = self.optimize_by_ort(m, use_external_data_format=use_external_data_format)
+ if float16:
+ logger.info("Convert to float16 ...")
+ m.convert_float_to_float16(
+ keep_io_types=keep_io_types,
+ op_block_list=fp32_op_list,
+ )
+
m.get_operator_statistics()
m.get_fused_operator_statistics()
m.save_model_to_file(optimized_onnx_path, use_external_data_format=use_external_data_format)
diff --git a/onnxruntime/python/tools/transformers/models/stable_diffusion/pipeline_txt2img.py b/onnxruntime/python/tools/transformers/models/stable_diffusion/pipeline_txt2img.py
index 444b6d9a8c..b9759b44e7 100644
--- a/onnxruntime/python/tools/transformers/models/stable_diffusion/pipeline_txt2img.py
+++ b/onnxruntime/python/tools/transformers/models/stable_diffusion/pipeline_txt2img.py
@@ -79,7 +79,7 @@ class Txt2ImgPipeline(StableDiffusionPipeline):
latents = self.denoise_latent(latents, text_embeddings, guidance=guidance)
# VAE decode latent
- images = self.decode_latent(latents)
+ images = self.decode_latent(latents / self.vae_scaling_factor)
torch.cuda.synchronize()
e2e_toc = time.perf_counter()
diff --git a/onnxruntime/python/tools/transformers/onnx_model.py b/onnxruntime/python/tools/transformers/onnx_model.py
index 392f2f9489..5fda3e6d84 100644
--- a/onnxruntime/python/tools/transformers/onnx_model.py
+++ b/onnxruntime/python/tools/transformers/onnx_model.py
@@ -1126,7 +1126,9 @@ class OnnxModel:
op = (node.domain + ":" if include_domain and node.domain else "") + node.op_type
op_count[op] = 1 if op not in op_count else (op_count[op] + 1)
- logger.info(f"Operators:{op_count}")
+ # Sorted by count in the descending order, then by key in alphabetical order.
+ logger.info(f"Operators:{sorted(op_count.items(), key=lambda kv:(-kv[1], kv[0]))}")
+
return op_count
@staticmethod
diff --git a/onnxruntime/python/tools/transformers/onnx_model_bert.py b/onnxruntime/python/tools/transformers/onnx_model_bert.py
index 7a69922e67..882100a0d0 100644
--- a/onnxruntime/python/tools/transformers/onnx_model_bert.py
+++ b/onnxruntime/python/tools/transformers/onnx_model_bert.py
@@ -488,16 +488,22 @@ class BertOnnxModel(OnnxModel):
logger.info(f"Optimized operators: {op_count}")
return op_count
- def is_fully_optimized(self):
+ def is_fully_optimized(self, fused_op_count=None):
"""
Returns True when the model is fully optimized.
"""
- op_count = self.get_fused_operator_statistics()
- embed = op_count["EmbedLayerNormalization"]
- attention = op_count["Attention"] + op_count["MultiHeadAttention"] + op_count["QOrderedAttention"]
- gelu = op_count["Gelu"] + op_count["BiasGelu"] + op_count["FastGelu"]
- layer_norm = op_count["LayerNormalization"] + op_count["SkipLayerNormalization"]
- simple_layer_norm = op_count["SimplifiedLayerNormalization"] + op_count["SkipSimplifiedLayerNormalization"]
+ if fused_op_count is None:
+ fused_op_count = self.get_fused_operator_statistics()
+
+ def op_count(op_name: str):
+ return fused_op_count.get(op_name) or 0
+
+ embed = op_count("EmbedLayerNormalization")
+ attention = op_count("Attention") + op_count("MultiHeadAttention") + op_count("QOrderedAttention")
+ gelu = op_count("Gelu") + op_count("BiasGelu") + op_count("FastGelu")
+ layer_norm = op_count("LayerNormalization") + op_count("SkipLayerNormalization")
+ simple_layer_norm = op_count("SimplifiedLayerNormalization") + op_count("SkipSimplifiedLayerNormalization")
+
is_perfect = (
(embed > 0)
and (attention > 0)
@@ -512,13 +518,13 @@ class BertOnnxModel(OnnxModel):
logger.debug("Simple Layer Normalization not fused")
if gelu == 0:
- logger.debug("Gelu/FastGelu not fused")
+ logger.debug("Gelu (or FastGelu) not fused")
if embed == 0:
- logger.debug("Embed Layer not fused")
+ logger.debug("EmbedLayerNormalization not fused")
if attention == 0:
- logger.warning("Attention not fused")
+ logger.warning("Attention (or MultiHeadAttention) not fused")
return is_perfect
diff --git a/onnxruntime/python/tools/transformers/onnx_model_unet.py b/onnxruntime/python/tools/transformers/onnx_model_unet.py
index 294641dd1e..4d15b9288e 100644
--- a/onnxruntime/python/tools/transformers/onnx_model_unet.py
+++ b/onnxruntime/python/tools/transformers/onnx_model_unet.py
@@ -12,6 +12,7 @@ from fusion_biassplitgelu import FusionBiasSplitGelu
from fusion_group_norm import FusionGroupNorm
from fusion_nhwc_conv import FusionNhwcConv
from fusion_options import FusionOptions
+from fusion_skip_group_norm import FusionSkipGroupNorm
from fusion_transpose import FusionInsertTranspose, FusionTranspose
from onnx import ModelProto
from onnx_model import OnnxModel
@@ -57,8 +58,8 @@ class UnetOnnxModel(BertOnnxModel):
logger.info("Removed %d Div nodes", len(nodes_to_remove))
def convert_conv_to_nhwc(self):
- # Do not update weight here since save external data has a bug
- conv_to_nhwc_conv = FusionNhwcConv(self, update_weight=False)
+ # Transpose weights in offline might help since ORT does not apply constant-folding on Transpose nodes.
+ conv_to_nhwc_conv = FusionNhwcConv(self, update_weight=True)
conv_to_nhwc_conv.apply()
def merge_adjacent_transpose(self):
@@ -150,6 +151,10 @@ class UnetOnnxModel(BertOnnxModel):
# Remove reshape nodes that having same shape of input and output based on symbolic shape inference.
self.utils.remove_useless_reshape_nodes()
+ if (options is None) or options.enable_skip_group_norm:
+ skip_group_norm_fusion = FusionSkipGroupNorm(self)
+ skip_group_norm_fusion.apply()
+
if (options is None) or options.enable_bias_skip_layer_norm:
# Fuse SkipLayerNormalization and Add Bias before it.
self.fuse_add_bias_skip_layer_norm()
@@ -181,6 +186,7 @@ class UnetOnnxModel(BertOnnxModel):
"SkipLayerNormalization",
"BiasSplitGelu",
"GroupNorm",
+ "SkipGroupNorm",
"NhwcConv",
"BiasAdd",
]
diff --git a/onnxruntime/python/tools/transformers/onnx_model_vae.py b/onnxruntime/python/tools/transformers/onnx_model_vae.py
index 9e79014e71..de8b59074a 100644
--- a/onnxruntime/python/tools/transformers/onnx_model_vae.py
+++ b/onnxruntime/python/tools/transformers/onnx_model_vae.py
@@ -32,6 +32,7 @@ class VaeOnnxModel(UnetOnnxModel):
ops = [
"Attention",
"GroupNorm",
+ "SkipGroupNorm",
"NhwcConv",
]
for op in ops:
diff --git a/onnxruntime/python/tools/transformers/optimizer.py b/onnxruntime/python/tools/transformers/optimizer.py
index f47bbaffaa..b2d6423a45 100644
--- a/onnxruntime/python/tools/transformers/optimizer.py
+++ b/onnxruntime/python/tools/transformers/optimizer.py
@@ -510,11 +510,14 @@ def main():
if args.input_int32:
optimizer.change_graph_inputs_to_int32()
- if args.model_type in set(MODEL_TYPES.keys()):
- if optimizer.is_fully_optimized():
- logger.info("The model has been fully optimized.")
- else:
- logger.info("The model has been optimized.")
+ # Print the operator statistics might help end user.
+ optimizer.get_operator_statistics()
+
+ fused_op_count = optimizer.get_fused_operator_statistics()
+ if "bert" in args.model_type and optimizer.is_fully_optimized(fused_op_count):
+ logger.info("The model has been fully optimized.")
+ else:
+ logger.info("The model has been optimized.")
if args.convert_to_packing_mode:
if args.model_type == "bert":
diff --git a/onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc b/onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc
index 29d8219c16..55f01bf0d3 100644
--- a/onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc
+++ b/onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc
@@ -25,7 +25,8 @@ static void RunTest(
int64_t interleaved,
bool use_float16,
bool disable_cpu,
- bool disable_cuda) {
+ bool disable_cuda,
+ bool disable_dml) {
// input : (batch_size, sequence_length, hidden_size)
// position ids : (1) or (batch_size, sequence_length)
// cos cache : (max_sequence_length, head_size / 2)
@@ -50,9 +51,14 @@ static void RunTest(
int min_cuda_architecture = use_float16 ? 530 : 0;
bool enable_cuda = HasCudaEnvironment(min_cuda_architecture);
+ bool enable_dml = (nullptr != DefaultDmlExecutionProvider().get()) && !disable_dml;
+
if (enable_cuda && !disable_cuda) {
execution_providers.push_back(DefaultCudaExecutionProvider());
}
+ if (enable_dml && !disable_dml) {
+ execution_providers.push_back(DefaultDmlExecutionProvider());
+ }
if (!use_float16 && !disable_cpu) {
execution_providers.push_back(DefaultCpuExecutionProvider());
}
@@ -107,9 +113,10 @@ static void RunTests(const std::vector& input_data,
interleaved,
false, /* use_fp16 */
false, /* disable_cpu */
- true /* disable_cuda */);
+ true, /* disable_cuda */
+ true /* disable_dml */);
- // FP32 test for CUDA
+ // FP32 test for CUDA and DML
RunTest(input_data,
position_ids,
cos_cache,
@@ -123,9 +130,10 @@ static void RunTests(const std::vector& input_data,
interleaved,
false, /* use_fp16 */
false, /* disable_cpu */
- false /* disable_cuda */);
+ false, /* disable_cuda */
+ false /* disable_dml */);
- // FP16 test for CUDA
+ // FP16 test for CUDA and DML
if (use_float16) {
RunTest(input_data,
position_ids,
@@ -138,9 +146,10 @@ static void RunTests(const std::vector& input_data,
num_heads,
max_sequence_length,
interleaved,
- true, /* use_fp16 */
- true, /* disable_cpu */
- false /* disable_cuda*/);
+ true, /* use_fp16 */
+ true, /* disable_cpu */
+ false, /* disable_cuda*/
+ false /* disable_dml */);
}
}