From 8f063305aacd23aadf77cda4c52e0592a02a7161 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Tue, 7 Nov 2023 17:45:30 -0800 Subject: [PATCH] Cherry pick LLaMA or SDXL to 1.16.2 release (round 3) (#18323) --- docs/OperatorKernels.md | 1 + .../cuda/bert/skip_layer_norm_impl.cu | 7 +- .../Operators/DmlOperatorRotaryEmbedding.cpp | 436 ++++++++++++++++++ .../src/Operators/OperatorRegistration.cpp | 4 + .../dml/OperatorAuthorHelper/Attributes.h | 1 + .../dml/OperatorAuthorHelper/OperatorHelper.h | 1 + .../OperatorAuthorHelper/OperatorVersions.h | 1 + .../python/onnxruntime_pybind_iobinding.cc | 2 - .../python/tools/symbolic_shape_infer.py | 2 + .../python/tools/transformers/float16.py | 3 +- .../tools/transformers/fusion_group_norm.py | 22 +- .../tools/transformers/fusion_options.py | 11 + .../transformers/fusion_rotary_attention.py | 10 +- .../transformers/fusion_skip_group_norm.py | 255 ++++++++++ .../models/stable_diffusion/demo_txt2img.py | 3 + .../stable_diffusion/demo_txt2img_xl.py | 8 + .../models/stable_diffusion/demo_utils.py | 5 +- .../stable_diffusion/diffusion_models.py | 36 +- .../stable_diffusion/diffusion_schedulers.py | 1 + .../models/stable_diffusion/engine_builder.py | 21 +- .../engine_builder_ort_cuda.py | 41 +- .../models/stable_diffusion/ort_optimizer.py | 50 +- .../stable_diffusion/pipeline_txt2img.py | 2 +- .../python/tools/transformers/onnx_model.py | 4 +- .../tools/transformers/onnx_model_bert.py | 26 +- .../tools/transformers/onnx_model_unet.py | 10 +- .../tools/transformers/onnx_model_vae.py | 1 + .../python/tools/transformers/optimizer.py | 13 +- .../contrib_ops/rotary_embedding_op_test.cc | 25 +- 29 files changed, 915 insertions(+), 87 deletions(-) create mode 100644 onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRotaryEmbedding.cpp create mode 100644 onnxruntime/python/tools/transformers/fusion_skip_group_norm.py 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 */); } }