diff --git a/.github/stale.yml b/.github/stale.yml new file mode 100644 index 0000000000..fcff8d38e6 --- /dev/null +++ b/.github/stale.yml @@ -0,0 +1,22 @@ +# Number of days of inactivity before an issue becomes stale +daysUntilStale: 60 + +# Number of days of inactivity before a stale issue is closed +daysUntilClose: 7 + +# Issues with these labels will never be considered stale +exemptLabels: + - "contributions are welcome" + - documentation + - enhancement + +# Label to use when marking an issue as stale +staleLabel: wontfix + +# Comment to post when marking an issue as stale. Set to `false` to disable +markComment: > + This issue has been automatically marked as stale due to inactivity and will be closed in 7 days if no further activity occurs. If further support is needed, please provide an update and/or more details. + +# Comment to post when closing a stale issue. Set to `false` to disable +closeComment: > + This issue has been automatically closed due to inactivity. Please reactivate if further support is needed. diff --git a/cmake/CMakeLists.txt b/cmake/CMakeLists.txt index 11c51fb763..f807aaeb9c 100644 --- a/cmake/CMakeLists.txt +++ b/cmake/CMakeLists.txt @@ -834,9 +834,8 @@ if (onnxruntime_USE_CUDA) string(APPEND CMAKE_CUDA_FLAGS "-cudart shared") endif() enable_language(CUDA) - string(REGEX REPLACE "([0-9]+)\\.([0-9]+).*" "\\1" CUDA_VERSION_MAJOR "${CMAKE_CUDA_COMPILER_VERSION}") - message( STATUS "CUDA_VERSION_MAJOR: ${CUDA_VERSION_MAJOR}") - if (CUDA_VERSION_MAJOR EQUAL 11) + message( STATUS "CMAKE_CUDA_COMPILER_VERSION: ${CMAKE_CUDA_COMPILER_VERSION}") + if (CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL 11) set(CMAKE_CUDA_STANDARD 14) else() set(CMAKE_CUDA_STANDARD 11) @@ -867,13 +866,16 @@ if (onnxruntime_USE_CUDA) list(APPEND onnxruntime_EXTERNAL_LIBRARIES ${ONNXRUNTIME_CUDA_LIBRARIES}) # the following compute capabilities are deprecated in CUDA 11 Toolkit - if (CUDA_VERSION_MAJOR LESS 11) + if (CMAKE_CUDA_COMPILER_VERSION VERSION_LESS 11) set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_30,code=sm_30") # K series set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_50,code=sm_50") # M series endif() set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_52,code=sm_52") # M60 set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_60,code=sm_60") # P series set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_70,code=sm_70") # V series + if (CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL 11) + set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_80,code=sm_80") # A series + endif() set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} --expt-relaxed-constexpr --default-stream legacy") if (NOT WIN32) set(CUDA_NVCC_FLAGS "${CUDA_NVCC_FLAGS} --expt-relaxed-constexpr --compiler-options -fPIC") diff --git a/cmake/onnxruntime_providers.cmake b/cmake/onnxruntime_providers.cmake index 4d9fc96061..ea3c9bbeeb 100644 --- a/cmake/onnxruntime_providers.cmake +++ b/cmake/onnxruntime_providers.cmake @@ -246,7 +246,7 @@ if (onnxruntime_USE_CUDA) set_target_properties(onnxruntime_providers_cuda PROPERTIES LINKER_LANGUAGE CUDA) set_target_properties(onnxruntime_providers_cuda PROPERTIES FOLDER "ONNXRuntime") - if (CUDA_VERSION_MAJOR LESS 11) + if (CMAKE_CUDA_COMPILER_VERSION VERSION_LESS 11) target_include_directories(onnxruntime_providers_cuda PRIVATE ${PROJECT_SOURCE_DIR}/external/cub) endif() diff --git a/onnxruntime/contrib_ops/cpu/bert/attention.cc b/onnxruntime/contrib_ops/cpu/bert/attention.cc index 760e86a1a9..8fc118c7e7 100644 --- a/onnxruntime/contrib_ops/cpu/bert/attention.cc +++ b/onnxruntime/contrib_ops/cpu/bert/attention.cc @@ -48,6 +48,7 @@ Status AttentionBase::CheckInputs(const Tensor* input, dims.size()); } int batch_size = static_cast(dims[0]); + int sequence_length = static_cast(dims[1]); int hidden_size = static_cast(dims[2]); if (hidden_size % num_heads_ != 0) { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, @@ -74,25 +75,10 @@ Status AttentionBase::CheckInputs(const Tensor* input, } if (bias_dims[0] != weights_dims[1]) { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, - "Input 2 dimension 0 should have same length as dimension 1 of input 1"); - } - - if (mask_index != nullptr) { // mask_index is optional - // unidirectional (like GPT2) does not need mask input. Here we do not allowed the input for unidirectional. - if (is_unidirectional_) { - return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'mask_index' is not allowed for unidirectional"); - } - - const auto& mask_dims = mask_index->Shape().GetDims(); - if (mask_dims.size() != 1) { - return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'mask_index' is expected to have 1 dimension, got ", - mask_dims.size()); - } - if (static_cast(mask_dims[0]) != batch_size) { - return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Inputs 'mask_index' and 'input' shall have same length at dimension 0"); - } + "Input 'bias' dimension 0 should have same length as dimension 1 of input 'weights'"); } + int past_sequence_length = 0; if (past != nullptr) { // past is optional if (!is_unidirectional_) { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'past' is only allowed for unidirectional"); @@ -115,8 +101,24 @@ Status AttentionBase::CheckInputs(const Tensor* input, if (static_cast(past_dims[4]) != hidden_size / num_heads_) { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Inputs 'past' dimension 2 shall have length of ", hidden_size / num_heads_); } + past_sequence_length = static_cast(past_dims[3]); } + if (mask_index != nullptr) { // mask_index is optional + const auto& mask_dims = mask_index->Shape().GetDims(); + if (mask_dims.size() == 1) { + if (static_cast(mask_dims[0]) != batch_size && static_cast(mask_dims[0]) != 2 * batch_size) { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Inputs 'mask_index' dimension 0 shall have length of batch_size or 2 * batch_size"); + } + } else if (mask_dims.size() == 2) { + if (static_cast(mask_dims[0]) != batch_size || static_cast(mask_dims[1]) != past_sequence_length + sequence_length) { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Inputs 'mask_index' with raw attention mask shall have shape batch_size x (past_sequence_length + sequence_length)"); + } + } else { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'mask_index' is expected to have 1 or 2 dimensions, got ", + mask_dims.size()); + } + } return Status::OK(); } @@ -174,7 +176,6 @@ Status Attention::Compute(OpKernelContext* context) const { ORT_RETURN_IF_ERROR(context->GetTempSpaceAllocator(&allocator)); auto* tp = context->GetOperatorThreadPool(); - // Compute Q, K, V // gemm_data(BS, 3NH) = input(BS, NH) x weights(NH, 3NH) + bias(3NH) auto gemm_data = allocator->Alloc(SafeInt(batch_size) * sequence_length * 3 * hidden_size * element_size); diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_cpu_base.h b/onnxruntime/contrib_ops/cpu/bert/attention_cpu_base.h index 1f01504ded..4252098dbf 100644 --- a/onnxruntime/contrib_ops/cpu/bert/attention_cpu_base.h +++ b/onnxruntime/contrib_ops/cpu/bert/attention_cpu_base.h @@ -48,26 +48,21 @@ class AttentionCPUBase : public AttentionBase { auto attention_probs = allocator->Alloc(attention_probs_bytes); BufferUniquePtr scratch_buffer(attention_probs, BufferDeleter(allocator)); - size_t mask_data_bytes = 0; - if (mask_index != nullptr) { - mask_data_bytes = SafeInt(batch_size) * sequence_length * all_sequence_length * sizeof(T); - } else if (is_unidirectional_) { - mask_data_bytes = SafeInt(sequence_length) * all_sequence_length * sizeof(T); - } - void* mask_data = nullptr; - if (mask_data_bytes > 0) { + if (mask_index != nullptr || (is_unidirectional_ && sequence_length > 1)) { + size_t mask_data_bytes = SafeInt(batch_size) * sequence_length * all_sequence_length * sizeof(T); mask_data = allocator->Alloc(mask_data_bytes); memset(mask_data, 0, mask_data_bytes); } BufferUniquePtr mask_data_buffer(mask_data, BufferDeleter(allocator)); const int32_t* mask_index_data = mask_index != nullptr ? mask_index->template Data() : nullptr; + const std::vector* mask_index_dims = mask_index != nullptr ? &(mask_index->Shape().GetDims()) : nullptr; const T* past_data = past != nullptr ? past->template Data() : nullptr; T* present_data = present != nullptr ? present->template MutableData() : nullptr; ComputeAttentionProbs(static_cast(attention_probs), Q, K, - mask_index_data, static_cast(mask_data), + mask_index_data, mask_index_dims, static_cast(mask_data), batch_size, sequence_length, past_sequence_length, head_size, past_data, present_data, tp); @@ -89,17 +84,18 @@ class AttentionCPUBase : public AttentionBase { // 1 x mask_data(B, N, S, S*) // II.attention_probs(B, N, S, S*) = Softmax(attention_probs) template - void ComputeAttentionProbs(T* attention_probs, // output buffer for the attention probs. Its size is BxNxSxS - const T* Q, // Q data. Its size is BxNxSxH - const T* K, // k data. Its size is BxNxSxH - const int32_t* mask_index, // mask index. nullptr if no mask or its size is B - T* mask_data, // buffer for mask data. Its size is: SxS* if is_unidirectional_; BxSxS* if mask_index; null otherwise - int batch_size, // batch size of self-attention - int sequence_length, // sequence length of self-attention - int past_sequence_length, // sequence length of past state - int head_size, // head size of self-attention - const T* past, // past state - T* present, // present state + void ComputeAttentionProbs(T* attention_probs, // output buffer for the attention probs. Its size is BxNxSxS + const T* Q, // Q data. Its size is BxNxSxH + const T* K, // k data. Its size is BxNxSxH + const int32_t* mask_index, // mask index. nullptr if no mask or its size is B + const std::vector* mask_index_dims, // mask index shape + T* mask_data, // buffer for mask data. Its size is: SxS* if is_unidirectional_; BxSxS* if mask_index; null otherwise + int batch_size, // batch size of self-attention + int sequence_length, // sequence length of self-attention + int past_sequence_length, // sequence length of past state + int head_size, // head size of self-attention + const T* past, // past state + T* present, // present state ThreadPool* tp) const { const int all_sequence_length = past_sequence_length + sequence_length; // S* = S' + S const size_t past_chunk_length = static_cast(past_sequence_length * head_size); // S' x H @@ -108,7 +104,7 @@ class AttentionCPUBase : public AttentionBase { { if (mask_data != nullptr) { - PrepareMask(mask_index, mask_data, is_unidirectional_, batch_size, sequence_length, past_sequence_length); + PrepareMask(mask_index, mask_index_dims, mask_data, is_unidirectional_, batch_size, sequence_length, past_sequence_length); } else { // no any mask memset(attention_probs, 0, batch_size * num_heads_ * sequence_length * all_sequence_length * sizeof(T)); } @@ -123,9 +119,9 @@ class AttentionCPUBase : public AttentionBase { for (std::ptrdiff_t i = begin; i != end; ++i) { const std::ptrdiff_t batch_index = i / num_heads_; - // broadcast mask data: SxS* or (Bx)SxS* -> (BxNx)SxS* + // broadcast mask data: (Bx)SxS* -> (BxNx)SxS* if (mask_data != nullptr) { - const T* broadcast_data_src = is_unidirectional_ ? reinterpret_cast(mask_data) : reinterpret_cast(mask_data) + batch_index * sequence_length * all_sequence_length; + const T* broadcast_data_src = reinterpret_cast(mask_data) + batch_index * sequence_length * all_sequence_length; T* broadcast_data_dest = reinterpret_cast(attention_probs) + sequence_length * all_sequence_length * i; memcpy(broadcast_data_dest, broadcast_data_src, sequence_length * all_sequence_length * sizeof(T)); } diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_helper.h b/onnxruntime/contrib_ops/cpu/bert/attention_helper.h index 7507cb2ab2..5acb8c2a30 100644 --- a/onnxruntime/contrib_ops/cpu/bert/attention_helper.h +++ b/onnxruntime/contrib_ops/cpu/bert/attention_helper.h @@ -61,38 +61,66 @@ inline void ComputeAttentionSoftmaxInplace(float* score, int N, int D, ThreadPoo template void PrepareMask(const int32_t* mask_index, + const std::vector* mask_index_dims, T* mask_data, bool is_unidirectional, int batch_size, int sequence_length, int past_sequence_length) { const int all_sequence_length = past_sequence_length + sequence_length; - T* p_mask = mask_data; - if (is_unidirectional) { - // unidirectional mask has shape SxS* - for (int s_i = 0; s_i < sequence_length - 1; s_i++) { - for (int m_i = past_sequence_length + s_i + 1; m_i < all_sequence_length; m_i++) { - p_mask[s_i * all_sequence_length + m_i] = static_cast(-10000.0); - } - } - return; - } - ORT_ENFORCE(mask_index, "mask index should not be null."); + // mask_data has been filled with 0, and its shape is BxSxS* + T* p_mask = mask_data; + + bool is_raw_attention_mask = (nullptr != mask_index_dims && mask_index_dims->size() == 2); + bool has_mask_start_position = (nullptr != mask_index_dims && mask_index_dims->size() == 1 && static_cast(mask_index_dims->at(0)) == 2 * batch_size); + for (int b_i = 0; b_i < batch_size; b_i++) { // TODO: mask_index can be used in softmax to save some calculation. - // Convert mask_index to mask (-10000 means out of range, which will be 0 after softmax): B => BxS* - int valid_length = mask_index[b_i]; - for (int m_i = valid_length; m_i < all_sequence_length; m_i++) { - p_mask[m_i] = static_cast(-10000.0); + + if (nullptr != mask_index) { + if (is_raw_attention_mask) { + // Raw attention mask has value 0 or 1. Here we convert 0 to -10000.0, and 1 to 0.0. + const int32_t* raw_mask = mask_index + b_i * all_sequence_length; + for (int m_i = 0; m_i < all_sequence_length; m_i++) { + p_mask[m_i] = (raw_mask[m_i] > 0) ? static_cast(0.0f) : static_cast(-10000.0f); + } + } else { + // mask_index is 1D: (B) or (2B) => (Bx)S* + + // Handle right-side padding: mask value at or after the end position will be -10000.0 + int end_position = mask_index[b_i]; + for (int m_i = end_position; m_i < all_sequence_length; m_i++) { + p_mask[m_i] = static_cast(-10000.0f); + } + + // Handle left-side padding: mask value before the start position will be -10000.0 + if (has_mask_start_position) { + int start_position = std::min(mask_index[b_i + batch_size], all_sequence_length); + for (int m_i = 0; m_i < start_position; m_i++) { + p_mask[m_i] = static_cast(-10000.0f); + } + } + } } - // Broadcast mask from BxS* to BxSxS* + // Broadcast mask from (Bx)S* to (Bx)SxS* for (int s_i = 1; s_i < sequence_length; s_i++) { memcpy(p_mask + s_i * all_sequence_length, p_mask, all_sequence_length * sizeof(T)); } - p_mask += sequence_length * sequence_length; + + // Apply unidirectional mask. + if (is_unidirectional) { + for (int s_i = 0; s_i < sequence_length - 1; s_i++) { + for (int m_i = past_sequence_length + s_i + 1; m_i < all_sequence_length; m_i++) { + p_mask[s_i * all_sequence_length + m_i] += static_cast(-10000.0f); + } + } + } + + p_mask += sequence_length * all_sequence_length; } + } // Concatenate a past state chunk S'xH with input state chunk SxH into present state chunk S*xH diff --git a/onnxruntime/contrib_ops/cuda/bert/attention.cc b/onnxruntime/contrib_ops/cuda/bert/attention.cc index 158d0e8cc0..a6b3f0b3b0 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention.cc +++ b/onnxruntime/contrib_ops/cuda/bert/attention.cc @@ -89,6 +89,7 @@ Status Attention::ComputeInternal(OpKernelContext* context) const { if (!LaunchAttentionKernel( reinterpret_cast(gemm_buffer.get()), nullptr == mask_index ? nullptr : mask_index->template Data(), + nullptr == mask_index ? nullptr : &(mask_index->Shape().GetDims()), output->template MutableData(), batch_size, sequence_length, @@ -100,8 +101,7 @@ Status Attention::ComputeInternal(OpKernelContext* context) const { is_unidirectional_, past_sequence_length, nullptr == past ? nullptr : past->template Data(), - nullptr == present ? nullptr : present->template MutableData() - )) { + nullptr == present ? nullptr : present->template MutableData())) { // Get last error to reset it to cudaSuccess. CUDA_CALL(cudaGetLastError()); return Status(common::ONNXRUNTIME, common::FAIL); diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu index 9e7a61b518..54206eb5fc 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu +++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu @@ -40,8 +40,8 @@ static size_t AlignTo(size_t a, size_t b) { return CeilDiv(a, b) * b; } -size_t ScratchSize(size_t element_size, int batch_size, int num_heads, int sequence_length, int past_sequence_length) { - const size_t len = batch_size * num_heads * sequence_length * (sequence_length + past_sequence_length); +size_t ScratchSize(size_t element_size, int batch_size, int num_heads, int sequence_length, int all_sequence_length) { + const size_t len = batch_size * num_heads * sequence_length * all_sequence_length; const size_t bytes = len * element_size; const size_t alignment = 256; @@ -57,11 +57,16 @@ size_t GetAttentionWorkspaceSize( int sequence_length, int past_sequence_length) { size_t qkv_size = 3 * batch_size * sequence_length * num_heads * head_size * element_size; - return qkv_size + 2 * ScratchSize(element_size, batch_size, num_heads, sequence_length, past_sequence_length); + return qkv_size + 2 * ScratchSize(element_size, batch_size, num_heads, sequence_length, past_sequence_length + sequence_length); } template -__device__ inline void Softmax(const int past_sequence_length, const int sequence_length, const int valid_length, const T* input, T* output, bool is_unidirectional) { +__device__ inline void Softmax(const int all_sequence_length, + const int sequence_length, + const int valid_end, + const int valid_start, + const T* input, + T* output) { using BlockReduce = cub::BlockReduce; __shared__ typename BlockReduce::TempStorage tmp_storage; @@ -70,18 +75,17 @@ __device__ inline void Softmax(const int past_sequence_length, const int sequenc float thread_data_max(-CUDART_INF_F); - const int num_valid = is_unidirectional ? past_sequence_length + (blockIdx.x % sequence_length) + 1 : valid_length; - // e^x is represented as infinity if x is large enough, like 100.f. // Infinity divided by Infinity is a NAN. Thus, softmax gets a NAN if one or more item are large enough. // a math transform as below is leveraged to get a stable softmax: // e^xi/(e^x1 + ...e^xn) = e^(xi - max) / (e^(x1 - max) + ... + e^(xn - max)) - const int all_sequence_length = past_sequence_length + sequence_length; const int offset = (blockIdx.y * gridDim.x + blockIdx.x) * all_sequence_length; - for (int i = threadIdx.x; i < num_valid; i += TPB) { - const int index = offset + i; - if (thread_data_max < float(input[index])) { - thread_data_max = float(input[index]); + for (int i = threadIdx.x; i < valid_end; i += TPB) { + if (i >= valid_start) { + const int index = offset + i; + if (thread_data_max < float(input[index])) { + thread_data_max = float(input[index]); + } } } @@ -94,10 +98,12 @@ __device__ inline void Softmax(const int past_sequence_length, const int sequenc __syncthreads(); float thread_data_sum(0.f); - for (int i = threadIdx.x; i < num_valid; i += TPB) { - const int index = offset + i; - const float val = input[index]; - thread_data_sum += expf(val - max_block); + for (int i = threadIdx.x; i < valid_end; i += TPB) { + if (i >= valid_start) { + const int index = offset + i; + const float val = input[index]; + thread_data_sum += expf(val - max_block); + } } const auto sum = BlockReduce(tmp_storage).Reduce(thread_data_sum, cub::Sum()); @@ -108,13 +114,19 @@ __device__ inline void Softmax(const int past_sequence_length, const int sequenc for (int i = threadIdx.x; i < all_sequence_length; i += TPB) { const int index = offset + i; - const float val = (i < num_valid) ? expf(float(input[index]) - max_block) * sum_reverse_block : 0.f; + const float val = (i >= valid_start && i < valid_end) ? expf(float(input[index]) - max_block) * sum_reverse_block : 0.f; output[index] = T(val); } } template -__device__ inline void SoftmaxSmall(const int past_sequence_length, const int sequence_length, const int valid_length, const T* input, T* output, bool is_unidirectional) { +__device__ inline void SoftmaxSmall(const int all_sequence_length, + const int sequence_length, + const int valid_end, + const int valid_start, + const T* input, + T* output, + bool is_unidirectional) { using BlockReduce = cub::BlockReduce; __shared__ typename BlockReduce::TempStorage tmp_storage; @@ -122,22 +134,32 @@ __device__ inline void SoftmaxSmall(const int past_sequence_length, const int se __shared__ float max_block; // Input dimension is BxNxSxS*; blockIdx.y is batch index b; gridDim.x=N*S; blockIdx.x is index within N*S; - const int all_sequence_length = past_sequence_length + sequence_length; const int offset = (blockIdx.y * gridDim.x + blockIdx.x) * all_sequence_length; const int index = offset + threadIdx.x; - const int num_valid = is_unidirectional ? past_sequence_length + (blockIdx.x % sequence_length) + 1 : valid_length; + bool is_valid = false; // whether it has attention mask == 1. + + // Update end position for unidirectional. + int end = valid_end; + if (is_unidirectional) { + int end_unid = all_sequence_length - sequence_length + (blockIdx.x % sequence_length) + 1; + if (end_unid <= valid_start) { + // In this situation, mask of [0, end_unid) and [valid_start, valid_end) has -10000, and [end_unid, valid_start) and [valid_end, all_seq_len) has -20000. + // So [0, end_unid) will also have value after softmax. + is_valid = threadIdx.x < end_unid; + } else { + end = min(valid_end, end_unid); + } + } + + is_valid = is_valid || (threadIdx.x >= valid_start && threadIdx.x < end); // e^x is represented as infinity if x is large enough, like 100.f. // Infinity divided by Infinity is a NAN. Thus, softmax gets a NAN if one or more item are large enough. // a math transform as below is leveraged to get a stable softmax: // e^xi/(e^x1 + ...e^xn) = e^(xi - max) / (e^(x1 - max) + ... + e^(xn - max)) - float thread_data_max(-CUDART_INF_F); - if (threadIdx.x < num_valid) { - thread_data_max = input[index]; - } - - const auto max = BlockReduce(tmp_storage).Reduce(thread_data_max, cub::Max(), num_valid); + float thread_data_max = is_valid ? float(input[index]) : float(-CUDART_INF_F); + const auto max = BlockReduce(tmp_storage).Reduce(thread_data_max, cub::Max(), end); // Store max value if (threadIdx.x == 0) { @@ -146,106 +168,239 @@ __device__ inline void SoftmaxSmall(const int past_sequence_length, const int se __syncthreads(); float thread_data_exp(0.f); - if (threadIdx.x < num_valid) { - const float val = input[index]; - thread_data_exp = expf(val - max_block); + if (is_valid) { + thread_data_exp = expf(float(input[index]) - max_block); } - const auto sum = BlockReduce(tmp_storage).Reduce(thread_data_exp, cub::Sum(), num_valid); + const auto sum = BlockReduce(tmp_storage).Reduce(thread_data_exp, cub::Sum(), end); - // Store max value + // Store value of 1.0/sum. if (threadIdx.x == 0) { - sum_reverse_block = (num_valid == 0) ? 0.f : (1.f) / sum; + sum_reverse_block = (1.f) / sum; } __syncthreads(); // threadIdx.x might be larger than all_sequence_length due to alignment to 32x. if (threadIdx.x < all_sequence_length) { - // this will be 0 for threadIdx.x >= num_valid output[index] = T(thread_data_exp * sum_reverse_block); } } template -__global__ void SoftmaxKernelSmall(const int past_sequence_length, const int sequence_length, const T* input, T* output, bool is_unidirectional) { - SoftmaxSmall(past_sequence_length, sequence_length, sequence_length, input, output, is_unidirectional); +__device__ inline void SoftmaxWithMask2DSmall(const int all_sequence_length, + const int sequence_length, + const int* attention_mask, // 2D attention mask + const T* input, + T* output, + const bool is_unidirectional, + const float scalar) { + using BlockReduce = cub::BlockReduce; + __shared__ typename BlockReduce::TempStorage tmp_storage; + + __shared__ float sum_reverse_block; + __shared__ float max_block; + + // Input dimension is BxNxSxS*; blockIdx.y is batch index b; gridDim.x=N*S; blockIdx.x is index within N*S; + int index = (blockIdx.y * gridDim.x + blockIdx.x) * all_sequence_length + threadIdx.x; + + float thread_data = -CUDART_INF_F; + if (threadIdx.x < all_sequence_length) { + const int& mask = attention_mask[blockIdx.y * all_sequence_length + threadIdx.x]; + float mask_value = mask > 0 ? 0.0f : -10000.0f; + + if (is_unidirectional) { + int from_index = all_sequence_length - sequence_length + (blockIdx.x % sequence_length); // offset of from token in all sequence length. + if (threadIdx.x > from_index) { + mask_value += -10000.0f; + } + } + + thread_data = float(input[index]) * scalar + mask_value; + } + + const float max = BlockReduce(tmp_storage).Reduce(thread_data, cub::Max(), all_sequence_length); + + // Store max value + if (threadIdx.x == 0) { + max_block = max; + } + __syncthreads(); + + float thread_data_exp = threadIdx.x < all_sequence_length ? expf(thread_data - max_block) : 0.0f; + const auto sum = BlockReduce(tmp_storage).Reduce(thread_data_exp, cub::Sum(), all_sequence_length); + + // Store value of 1.0/sum + if (threadIdx.x == 0) { + sum_reverse_block = (1.f) / sum; + } + __syncthreads(); + + if (threadIdx.x < all_sequence_length) { + output[index] = T(thread_data_exp * sum_reverse_block); + } } template -__global__ void SoftmaxKernel(const int past_sequence_length, const int sequence_length, const T* input, T* output, bool is_unidirectional) { - Softmax(past_sequence_length, sequence_length, sequence_length, input, output, is_unidirectional); +__global__ void SoftmaxKernelSmall(const int all_sequence_length, const int sequence_length, const T* input, T* output, bool is_unidirectional) { + SoftmaxSmall(all_sequence_length, sequence_length, all_sequence_length, 0, input, output, is_unidirectional); +} + +template +__global__ void SoftmaxKernel(const int all_sequence_length, const int sequence_length, const T* input, T* output) { + Softmax(all_sequence_length, sequence_length, all_sequence_length, 0, input, output); } template bool ComputeSoftmax( - cudaStream_t stream, const int past_sequence_length, const int sequence_length, const int batch_size, const int num_heads, + cudaStream_t stream, const int all_sequence_length, const int sequence_length, const int batch_size, const int num_heads, const T* input, T* output, bool is_unidirectional) { const dim3 grid(sequence_length * num_heads, batch_size, 1); - if (sequence_length <= 32) { + if (all_sequence_length <= 32) { const int blockSize = 32; - SoftmaxKernelSmall<<>>(past_sequence_length, sequence_length, input, output, is_unidirectional); - } else if (sequence_length <= 128) { + SoftmaxKernelSmall<<>>(all_sequence_length, sequence_length, input, output, is_unidirectional); + } else if (all_sequence_length <= 64) { + const int blockSize = 64; + SoftmaxKernelSmall<<>>(all_sequence_length, sequence_length, input, output, is_unidirectional); + } else if (all_sequence_length <= 128) { const int blockSize = 128; - SoftmaxKernelSmall<<>>(past_sequence_length, sequence_length, input, output, is_unidirectional); - } else if (sequence_length == 384) { - const int blockSize = 384; - SoftmaxKernelSmall<<>>(past_sequence_length, sequence_length, input, output, is_unidirectional); - } else { + SoftmaxKernelSmall<<>>(all_sequence_length, sequence_length, input, output, is_unidirectional); + } else if (all_sequence_length <= 256) { const int blockSize = 256; - SoftmaxKernel<<>>(past_sequence_length, sequence_length, input, output, is_unidirectional); + SoftmaxKernelSmall<<>>(all_sequence_length, sequence_length, input, output, is_unidirectional); + } else if (all_sequence_length <= 512) { + const int blockSize = 512; + SoftmaxKernelSmall<<>>(all_sequence_length, sequence_length, input, output, is_unidirectional); + } else if (all_sequence_length <= 1024) { + const int blockSize = 1024; + SoftmaxKernelSmall<<>>(all_sequence_length, sequence_length, input, output, is_unidirectional); + } else if (!is_unidirectional) { + const int blockSize = 1024; + SoftmaxKernel<<>>(all_sequence_length, sequence_length, input, output); + } else { + ORT_THROW("Attention CUDA operator does not support unidirectional with total sequence length > 1024."); } return CUDA_CALL(cudaPeekAtLastError()); } template -__global__ void MaskedSoftmaxKernelSmall(const int sequence_length, const int* mask_index, const T* input, T* output) { - __shared__ int num_valid; +__global__ void MaskedSoftmaxKernelSmall(const int all_sequence_length, const int sequence_length, const int* mask_end, const int* mask_start, const T* input, T* output, bool is_unidirectional) { + __shared__ int start_position; + __shared__ int end_position; if (threadIdx.x == 0) { - num_valid = min(sequence_length, mask_index[blockIdx.y]); + const int batch = blockIdx.y; + start_position = mask_start != nullptr ? max(0, mask_start[batch]) : 0; + end_position = min(all_sequence_length, mask_end[batch]); + + // Attend to no word has same effect as attend to all words. This is added to get parity with CPU result. + if (start_position >= end_position) { + start_position = 0; + end_position = all_sequence_length; + } } __syncthreads(); - SoftmaxSmall(0, sequence_length, num_valid, input, output, false); + SoftmaxSmall(all_sequence_length, sequence_length, end_position, start_position, input, output, is_unidirectional); } template -__global__ void MaskedSoftmaxKernel(const int sequence_length, const int* mask_index, const T* input, T* output) { - __shared__ int num_valid; +__global__ void MaskedSoftmaxKernel(const int all_sequence_length, const int sequence_length, const int* mask_end, const int* mask_start, const T* input, T* output) { + __shared__ int start_position; + __shared__ int end_position; if (threadIdx.x == 0) { - num_valid = min(sequence_length, mask_index[blockIdx.y]); + const int batch = blockIdx.y; + start_position = mask_start != nullptr ? max(0, mask_start[batch]) : 0; + end_position = min(all_sequence_length, mask_end[batch]); + + // Attend to no word has same effect as attend to all words. This is added to get parity with CPU result. + if (start_position >= end_position) { + start_position = 0; + end_position = all_sequence_length; + } } __syncthreads(); - Softmax(0, sequence_length, num_valid, input, output, false); + Softmax(all_sequence_length, sequence_length, end_position, start_position, input, output); +} + +template +__global__ void SoftmaxWithMask2DSmallKernel(const int all_sequence_length, const int sequence_length, const int* attention_mask, const T* input, T* output, const bool is_unidirectional, const float scalar) { + SoftmaxWithMask2DSmall(all_sequence_length, sequence_length, attention_mask, input, output, is_unidirectional, scalar); } template -bool ComputeMaskedSoftmax(cudaStream_t stream, const int sequence_length, const int batch_size, const int num_heads, - const int* mask_index, const T* input, T* output) { - // Mask is of length batch_size and assumes the valid region is contiguous starting - // from the beginning of the sequence - +bool ComputeSoftmaxWithMask1D(cudaStream_t stream, const int all_sequence_length, const int sequence_length, const int batch_size, const int num_heads, + const int* mask_index, const int* mask_start, const T* input, T* output, const bool is_unidirectional) { const dim3 grid(sequence_length * num_heads, batch_size, 1); - if (sequence_length <= 32) { + if (all_sequence_length <= 32) { const int blockSize = 32; MaskedSoftmaxKernelSmall - <<>>(sequence_length, mask_index, input, output); - } else if (sequence_length <= 128) { + <<>>(all_sequence_length, sequence_length, mask_index, mask_start, input, output, is_unidirectional); + } else if (all_sequence_length <= 64) { + const int blockSize = 64; + MaskedSoftmaxKernelSmall + <<>>(all_sequence_length, sequence_length, mask_index, mask_start, input, output, is_unidirectional); + } else if (all_sequence_length <= 128) { const int blockSize = 128; MaskedSoftmaxKernelSmall - <<>>(sequence_length, mask_index, input, output); - } else if (sequence_length == 384) { - const int blockSize = 384; - MaskedSoftmaxKernelSmall - <<>>(sequence_length, mask_index, input, output); - } else { + <<>>(all_sequence_length, sequence_length, mask_index, mask_start, input, output, is_unidirectional); + } else if (all_sequence_length <= 256) { const int blockSize = 256; + MaskedSoftmaxKernelSmall + <<>>(all_sequence_length, sequence_length, mask_index, mask_start, input, output, is_unidirectional); + } else if (all_sequence_length <= 512) { + const int blockSize = 512; + MaskedSoftmaxKernelSmall + <<>>(all_sequence_length, sequence_length, mask_index, mask_start, input, output, is_unidirectional); + } else if (all_sequence_length <= 1024) { + const int blockSize = 1024; + MaskedSoftmaxKernelSmall + <<>>(all_sequence_length, sequence_length, mask_index, mask_start, input, output, is_unidirectional); + } else if (!is_unidirectional) { + const int blockSize = 1024; MaskedSoftmaxKernel - <<>>(sequence_length, mask_index, input, output); + <<>>(all_sequence_length, sequence_length, mask_index, mask_start, input, output); + } else { + ORT_THROW("Attention CUDA operator does not support unidirectional with total sequence length > 1024."); + } + + return CUDA_CALL(cudaPeekAtLastError()); +} + +template +bool ComputeSoftmaxWithMask2D(cudaStream_t stream, const int all_sequence_length, const int sequence_length, const int batch_size, const int num_heads, + const int* attention_mask, const T* input, T* output, const bool is_unidirectional, const float scalar) { + const dim3 grid(sequence_length * num_heads, batch_size, 1); + + if (all_sequence_length <= 32) { + const int blockSize = 32; + SoftmaxWithMask2DSmallKernel + <<>>(all_sequence_length, sequence_length, attention_mask, input, output, is_unidirectional, scalar); + } else if (all_sequence_length <= 64) { + const int blockSize = 64; + SoftmaxWithMask2DSmallKernel + <<>>(all_sequence_length, sequence_length, attention_mask, input, output, is_unidirectional, scalar); + } else if (all_sequence_length <= 128) { + const int blockSize = 128; + SoftmaxWithMask2DSmallKernel + <<>>(all_sequence_length, sequence_length, attention_mask, input, output, is_unidirectional, scalar); + } else if (all_sequence_length <= 256) { + const int blockSize = 256; + SoftmaxWithMask2DSmallKernel + <<>>(all_sequence_length, sequence_length, attention_mask, input, output, is_unidirectional, scalar); + } else if (all_sequence_length <= 512) { + const int blockSize = 512; + SoftmaxWithMask2DSmallKernel + <<>>(all_sequence_length, sequence_length, attention_mask, input, output, is_unidirectional, scalar); + } else if (all_sequence_length <= 1024) { + const int blockSize = 1024; + SoftmaxWithMask2DSmallKernel + <<>>(all_sequence_length, sequence_length, attention_mask, input, output, is_unidirectional, scalar); + } else { + ORT_THROW("Attention CUDA operator does not supported 2D attention mask with total sequence length > 1024."); } return CUDA_CALL(cudaPeekAtLastError()); @@ -389,7 +544,7 @@ __global__ void ConcatPastToPresent(const int sequence_length, const int n = threadIdx.y; const int s = blockIdx.x; const int b = blockIdx.y; - const int is_v = blockIdx.z; // 0 for k, 1 for v + const int is_v = blockIdx.z; // 0 for k, 1 for v const int all_sequence_length = gridDim.x; const int batch_size = gridDim.y; @@ -409,7 +564,7 @@ __global__ void ConcatPastToPresent(const int sequence_length, const int past_NSH = num_heads * past_SH; const int in_offset = b * past_NSH + n * past_SH + s * H + h + is_v * (past_NSH * batch_size); present[out_offset] = past[in_offset]; -} else if (s < all_sequence_length) { + } else if (s < all_sequence_length) { const int SH = sequence_length * H; const int NSH = num_heads * SH; const int in_offset = b * NSH + n * SH + (s - past_sequence_length) * H + h + is_v * (NSH * batch_size); @@ -418,7 +573,7 @@ __global__ void ConcatPastToPresent(const int sequence_length, } bool LaunchConcatPastToPresent(cudaStream_t stream, - const int past_sequence_length, + const int all_sequence_length, const int sequence_length, const int batch_size, const int head_size, @@ -426,13 +581,11 @@ bool LaunchConcatPastToPresent(cudaStream_t stream, const float* past, const float* k_v, float* present) { - const int all_sequence_length = past_sequence_length + sequence_length; const dim3 grid(all_sequence_length, batch_size, 2); if (0 == (head_size & 1)) { const dim3 block(head_size / 2, num_heads, 1); ConcatPastToPresent<<>>(sequence_length, reinterpret_cast(past), reinterpret_cast(k_v), reinterpret_cast(present)); - } else - { + } else { const dim3 block(head_size, num_heads, 1); ConcatPastToPresent<<>>(sequence_length, past, k_v, present); } @@ -440,7 +593,7 @@ bool LaunchConcatPastToPresent(cudaStream_t stream, } bool LaunchConcatPastToPresent(cudaStream_t stream, - const int past_sequence_length, + const int all_sequence_length, const int sequence_length, const int batch_size, const int head_size, @@ -448,14 +601,13 @@ bool LaunchConcatPastToPresent(cudaStream_t stream, const half* past, const half* k_v, half* present) { - const int all_sequence_length = past_sequence_length + sequence_length; const dim3 grid(all_sequence_length, batch_size, 2); if (0 == (head_size % 4)) { const dim3 block(head_size / 4, num_heads, 1); ConcatPastToPresent<<>>(sequence_length, reinterpret_cast(past), reinterpret_cast(k_v), reinterpret_cast(present)); } else if (0 == (head_size & 1)) { const dim3 block(head_size / 2, num_heads, 1); - ConcatPastToPresent<<>>(sequence_length, reinterpret_cast(past), reinterpret_cast(k_v), reinterpret_cast(present)); + ConcatPastToPresent<<>>(sequence_length, reinterpret_cast(past), reinterpret_cast(k_v), reinterpret_cast(present)); } else { // this should be an "odd" case. probably not worth catching it in the half2 kernel. const dim3 block(head_size, num_heads, 1); ConcatPastToPresent<<>>(sequence_length, past, k_v, present); @@ -486,9 +638,10 @@ bool QkvToContext( cublasHandle_t& cublas, cudaStream_t stream, const int batch_size, const int sequence_length, const int num_heads, const int head_size, const size_t element_size, const T* input, T* output, T* workspace, - const int* mask_index, + const int* mask_index, const std::vector* mask_index_dims, bool is_unidirectional, int past_sequence_length, const T* past, T* present) { - const size_t bytes = ScratchSize(element_size, batch_size, num_heads, sequence_length, past_sequence_length); + const int all_sequence_length = past_sequence_length + sequence_length; + const size_t bytes = ScratchSize(element_size, batch_size, num_heads, sequence_length, all_sequence_length); T* scratch1 = workspace; T* scratch2 = scratch1 + (bytes / element_size); T* scratch3 = scratch2 + (bytes / element_size); @@ -513,9 +666,9 @@ bool QkvToContext( // Concat past (2xBxNxS'xH) to present (2xBxNxS*xH): // past_k (BxNxS'xH) + k (BxNxSxH) => present_k (BxNxS*xH) // past_v (BxNxS'xH) + v (BxNxSxH) => present_v (BxNxS*xH) - const int present_size_per_batch = (past_sequence_length + sequence_length) * head_size; + const int present_size_per_batch = all_sequence_length * head_size; if (nullptr != present) { - if (!LaunchConcatPastToPresent(stream, past_sequence_length, sequence_length, batch_size, head_size, num_heads, past, k, present)) { + if (!LaunchConcatPastToPresent(stream, all_sequence_length, sequence_length, batch_size, head_size, num_heads, past, k, present)) { return false; } @@ -524,24 +677,33 @@ bool QkvToContext( v = present + batches * present_size_per_batch; } + bool use_2d_attention_mask = (nullptr != mask_index && nullptr != mask_index_dims && mask_index_dims->size() == 2); + // compute Q*K' (as K'*Q), scaled by 1/sqrt(H) and store in scratch1: BxNxSxS* // Q: BxNxSxH, K (present_k): BxNxS*xH, Q*K': BxNxSxS* const float rsqrt_head_size = 1.f / sqrt(static_cast(head_size)); - const int all_sequence_length = past_sequence_length + sequence_length; const int temp_matrix_size = sequence_length * all_sequence_length; + T alpha = (T)(use_2d_attention_mask ? 1.0f : rsqrt_head_size); if (!CUBLAS_CALL(CublasGemmStridedBatched( - cublas, CUBLAS_OP_T, CUBLAS_OP_N, all_sequence_length, sequence_length, head_size, rsqrt_head_size, k, head_size, present_size_per_batch, + cublas, CUBLAS_OP_T, CUBLAS_OP_N, all_sequence_length, sequence_length, head_size, alpha, k, head_size, present_size_per_batch, q, head_size, size_per_batch, 0.f, scratch1, all_sequence_length, temp_matrix_size, batches))) { return false; } // apply softmax and store result P to scratch2: BxNxSxS* - if (nullptr != mask_index) { - if (!ComputeMaskedSoftmax(stream, sequence_length, batch_size, num_heads, mask_index, scratch1, scratch2)) { + if (use_2d_attention_mask) { // 2d attention mask + if (!ComputeSoftmaxWithMask2D(stream, all_sequence_length, sequence_length, batch_size, num_heads, mask_index, scratch1, scratch2, is_unidirectional, rsqrt_head_size)) { return false; } - } else { - if (!ComputeSoftmax(stream, past_sequence_length, sequence_length, batch_size, num_heads, scratch1, scratch2, is_unidirectional)) { + } else if (nullptr != mask_index) { // 1d mask index + ORT_ENFORCE(nullptr != mask_index_dims && mask_index_dims->size() == 1); + // mask_index has 1D shape: either (batch_size) or (2*batch_size). Only the later one has start postions. + const int* mask_start = (mask_index_dims->at(0) > batch_size) ? mask_index + batch_size : nullptr; + if (!ComputeSoftmaxWithMask1D(stream, all_sequence_length, sequence_length, batch_size, num_heads, mask_index, mask_start, scratch1, scratch2, is_unidirectional)) { + return false; + } + } else { // no mask + if (!ComputeSoftmax(stream, all_sequence_length, sequence_length, batch_size, num_heads, scratch1, scratch2, is_unidirectional)) { return false; } } @@ -560,6 +722,7 @@ bool QkvToContext( bool LaunchAttentionKernel( const void* input, const int* mask_index, + const std::vector* mask_index_dims, void* output, const int batch_size, const int sequence_length, @@ -579,13 +742,13 @@ bool LaunchAttentionKernel( return QkvToContext(cublas, stream, batch_size, sequence_length, num_heads, head_size, element_size, reinterpret_cast(input), reinterpret_cast(output), reinterpret_cast(workspace), - mask_index, is_unidirectional, + mask_index, mask_index_dims, is_unidirectional, past_sequence_length, reinterpret_cast(past), reinterpret_cast(present)); } else { return QkvToContext(cublas, stream, batch_size, sequence_length, num_heads, head_size, element_size, reinterpret_cast(input), reinterpret_cast(output), reinterpret_cast(workspace), - mask_index, is_unidirectional, + mask_index, mask_index_dims, is_unidirectional, past_sequence_length, reinterpret_cast(past), reinterpret_cast(present)); } } diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.h b/onnxruntime/contrib_ops/cuda/bert/attention_impl.h index 6e58e73072..8a4ecffe4b 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.h +++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.h @@ -16,20 +16,21 @@ size_t GetAttentionWorkspaceSize( int past_sequence_length); bool LaunchAttentionKernel( - const void* input, // Input tensor - const int* mask_index, // Mask index (length of each sequence). NULL means no mask. - void* output, // Output tensor - int batch_size, // Batch size (B) - int sequence_length, // Sequence length (S) - int num_heads, // Number of attention heads (N) - int head_size, // Hidden layer size per head (H) - void* workspace, // Temporary buffer - cublasHandle_t& cublas, // Cublas handle - const size_t element_size, // Element size of input tensor - bool is_unidirectional, // Whether there is unidirecitonal mask. - int past_sequence_length, // Sequence length in past state - const void* past, // Past state input - void* present // Present state output + const void* input, // Input tensor + const int* mask_index, // Attention mask raw data or index (end position of each sequence, or end positions and start positions). NULL means no mask. + const std::vector* mask_index_dims, // Mask index shape + void* output, // Output tensor + int batch_size, // Batch size (B) + int sequence_length, // Sequence length (S) + int num_heads, // Number of attention heads (N) + int head_size, // Hidden layer size per head (H) + void* workspace, // Temporary buffer + cublasHandle_t& cublas, // Cublas handle + const size_t element_size, // Element size of input tensor + bool is_unidirectional, // Whether there is unidirecitonal mask. + int past_sequence_length, // Sequence length in past state + const void* past, // Past state input + void* present // Present state output ); } // namespace cuda diff --git a/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc b/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc index 974fdecf2f..08e1f6b3b4 100644 --- a/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc +++ b/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc @@ -173,6 +173,7 @@ Status QAttention::ComputeInternal(OpKernelContext* context) const { if (!LaunchAttentionKernel( reinterpret_cast(gemm_buffer.get()), nullptr == mask_index ? nullptr : mask_index->template Data(), + nullptr == mask_index ? nullptr : &(mask_index->Shape().GetDims()), output->template MutableData(), batch_size, sequence_length, diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index eb67c26e52..93b69f59e9 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -291,9 +291,14 @@ const char* contrib_ops_auto_pad_doc = void RegisterBertSchemas() { static const char* Attention_ver1_doc = R"DOC( -Multi-Head Self Attention that can be either unidirectional (like GPT2) or bidirectional (like BERT). -The mask_index input is optional. Unidirectional and mask_index input are mutually exclusive. When unidirectional is 1, the -mask_index shall not be provided.)DOC"; +Multi-Head Self Attention that can be either unidirectional (like GPT-2) or bidirectional (like BERT). +The mask_index input is optional. Besides raw attention mask with shape (batch_size, past_sequence_length + sequence_length), +we also support other two formats: When input has right-side padding, mask_index is one dimension with shape (batch_size), +where value of each element is the end position, or valid length of actual sequence excluding padding. When input has +left-side padding, mask_index has shape (2 * batch_size), where the values are the exclusive end positions followed by +the inclusive start positions. When unidirectional is 1, and each token only attend to previous tokens. For GPT-2, both past +and present state are optional. Present state could appear in output even when past state is not in input. +)DOC"; ONNX_CONTRIB_OPERATOR_SCHEMA(Attention) .SetDomain(kMSDomain) @@ -308,7 +313,7 @@ mask_index shall not be provided.)DOC"; .Input(0, "input", "3D input tensor with shape (batch_size, sequence_length, hidden_size), hidden_size = num_heads * head_size", "T") .Input(1, "weight", "2D input tensor with shape (hidden_size, 3 * hidden_size)", "T") .Input(2, "bias", "1D input tensor with shape (3 * hidden_size)", "T") - .Input(3, "mask_index", "Attention mask index with shape (batch_size).", "M", OpSchema::Optional) + .Input(3, "mask_index", "Attention mask with shape (batch_size, past_sequence_length + sequence_length), or index with shape (batch_size) or (2 * batch_size).", "M", OpSchema::Optional) .Input(4, "past", "past state for key and value with shape (2, batch_size, num_heads, past_sequence_length, head_size).", "T", OpSchema::Optional) .Output(0, "output", "3D output tensor with shape (batch_size, append_length, hidden_size)", "T") .Output(1, "present", "present state for key and value with shape (2, batch_size, num_heads, past_sequence_length + sequence_length, head_size)", "T", OpSchema::Optional) diff --git a/onnxruntime/core/optimizer/nchwc_transformer.cc b/onnxruntime/core/optimizer/nchwc_transformer.cc index 6da13740c1..f00709dcbc 100644 --- a/onnxruntime/core/optimizer/nchwc_transformer.cc +++ b/onnxruntime/core/optimizer/nchwc_transformer.cc @@ -640,7 +640,8 @@ void NchwcTransformerImpl::TransformConcat(Node& node) { } // After doing a Conv/Add fusion, there may be an activation node that could now -// be fused into the Conv node as well. +// be fused into the Conv node as well. Otherwise, this is an elementwise +// operation that can directly use the NCHWc input. void NchwcTransformerImpl::TransformActivation(Node& node) { auto& input_defs = node.MutableInputDefs(); @@ -943,7 +944,9 @@ void NchwcTransformerImpl::Transform(Node& node) { TransformBinary(node, false); } else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "Concat", {4, 11})) { TransformConcat(node); - } else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "Relu", {6})) { + } else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "Relu", {6}) || + graph_utils::IsSupportedOptypeVersionAndDomain(node, "Sigmoid", {6}) || + graph_utils::IsSupportedOptypeVersionAndDomain(node, "Tanh", {6})) { TransformActivation(node); } else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "BatchNormalization", {7, 9})) { TransformBatchNormalization(node); diff --git a/onnxruntime/core/providers/cpu/math/matmul_helper.h b/onnxruntime/core/providers/cpu/math/matmul_helper.h index fca73c2d99..2d8d6c14aa 100644 --- a/onnxruntime/core/providers/cpu/math/matmul_helper.h +++ b/onnxruntime/core/providers/cpu/math/matmul_helper.h @@ -29,7 +29,7 @@ class MatMulComputeHelper { // A: [M1, M2, ... K], B: [N, K]^T // A: [M1, M2, ... K], B: [1, ..., 1, K, N] // A: [M1, M2, ... K], B: [1, ..., 1, N, K]^T - if (!transa && left_num_dims >= 2 && right_num_dims >= 2 && + if (!transa && left_num_dims >= 2 && right_num_dims >= 2 && left_num_dims >= right_num_dims && right_shape.SizeToDimension(right_num_dims - 1) == right_shape[right_num_dims - 2]) { M_ = left_shape.SizeToDimension(left_num_dims - 1); K_ = left_shape[left_num_dims - 1]; diff --git a/onnxruntime/core/providers/cuda/nn/conv.cc b/onnxruntime/core/providers/cuda/nn/conv.cc index f7004ee2be..0cccf3b41d 100644 --- a/onnxruntime/core/providers/cuda/nn/conv.cc +++ b/onnxruntime/core/providers/cuda/nn/conv.cc @@ -95,10 +95,6 @@ Status Conv::ComputeInternal(OpKernelContext* context) const { Tensor* Y = context->Output(0, TensorShape(s_.y_dims)); y_data = reinterpret_cast(Y->template MutableData()); - // special case when there is a dim value of 0 in the shape. - if (Y->Shape().Size() == 0) - return Status::OK(); - std::vector x_dims_cudnn = x_dims; std::vector y_dims_cudnn = y_dims; if (rank < 2) { @@ -112,12 +108,21 @@ Status Conv::ComputeInternal(OpKernelContext* context) const { strides.push_back(1); dilations.push_back(1); } - ORT_RETURN_IF_ERROR(s_.x_tensor.Set(x_dims_cudnn, CudnnTensor::GetDataType())); - ORT_RETURN_IF_ERROR(s_.y_tensor.Set(y_dims_cudnn, CudnnTensor::GetDataType())); if (w_dims_changed) ORT_RETURN_IF_ERROR(s_.filter_desc.Set(w_dims, CudnnTensor::GetDataType())); + // Special case when there is a dim value of 0 in the shape. + // Return only after we have cached the following for subsequent runs : + // 1) `w_dims` in the `filter_desc` + // 2) `y_dims` in s_.y_dims + if (Y->Shape().Size() == 0) { + return Status::OK(); + } + + ORT_RETURN_IF_ERROR(s_.x_tensor.Set(x_dims_cudnn, CudnnTensor::GetDataType())); + ORT_RETURN_IF_ERROR(s_.y_tensor.Set(y_dims_cudnn, CudnnTensor::GetDataType())); + cudnnConvolutionMode_t mode = CUDNN_CROSS_CORRELATION; ORT_RETURN_IF_ERROR(s_.conv_desc.Set(kernel_shape.size(), pads, strides, dilations, mode, CudnnTensor::GetDataType())); diff --git a/onnxruntime/core/providers/cuda/nn/conv_transpose.cc b/onnxruntime/core/providers/cuda/nn/conv_transpose.cc index 3fe3a53f71..e258d8d685 100644 --- a/onnxruntime/core/providers/cuda/nn/conv_transpose.cc +++ b/onnxruntime/core/providers/cuda/nn/conv_transpose.cc @@ -67,7 +67,7 @@ Status ConvTranspose::DoConvTranspose(OpKernelContext* context, bool dynamic_ { std::lock_guard lock(s_.mutex); - // TODO: add a global cache if need to handle cases for multiple frames running simultaneuously with different batch_size + // TODO: add a global cache if need to handle cases for multiple frames running simultaneously with different batch_size bool input_dims_changed = (s_.last_x_dims != x_dims); bool w_dims_changed = (s_.last_w_dims != w_dims); if (input_dims_changed || w_dims_changed) { @@ -82,11 +82,6 @@ Status ConvTranspose::DoConvTranspose(OpKernelContext* context, bool dynamic_ ConvTransposeAttributes::Prepare p; ORT_RETURN_IF_ERROR(conv_transpose_attrs_.PrepareForCompute(context, has_bias, p, dynamic_padding)); - // Bail out early if one of the dimensions is zero. - if (p.Y->Shape().Size() == 0) { - return Status::OK(); - } - auto y_dims = p.Y->Shape().GetDims(); if (x_dimensions == 3) { y_dims.insert(y_dims.begin() + 2, 1); @@ -98,12 +93,20 @@ Status ConvTranspose::DoConvTranspose(OpKernelContext* context, bool dynamic_ } s_.y_dims = y_dims; - ORT_RETURN_IF_ERROR(s_.x_tensor.Set(x_dims, CudnnTensor::GetDataType())); - ORT_RETURN_IF_ERROR(s_.y_tensor.Set(y_dims, CudnnTensor::GetDataType())); - if (w_dims_changed) ORT_RETURN_IF_ERROR(s_.filter_desc.Set(w_dims, CudnnTensor::GetDataType())); + // Special case when there is a dim value of 0 in the shape. + // Return only after we have cached the following for subsequent runs : + // 1) `w_dims` in the `filter_desc` + // 2) `y_dims` in s_.y_dims + if (p.Y->Shape().Size() == 0) { + return Status::OK(); + } + + ORT_RETURN_IF_ERROR(s_.x_tensor.Set(x_dims, CudnnTensor::GetDataType())); + ORT_RETURN_IF_ERROR(s_.y_tensor.Set(y_dims, CudnnTensor::GetDataType())); + cudnnConvolutionMode_t mode = CUDNN_CROSS_CORRELATION; ORT_RETURN_IF_ERROR(s_.conv_desc.Set(p.kernel_shape.size(), p.pads, p.strides, p.dilations, mode, CudnnTensor::GetDataType())); @@ -155,42 +158,49 @@ Status ConvTranspose::DoConvTranspose(OpKernelContext* context, bool dynamic_ s_.algo = perf.algo; s_.workspace_bytes = perf.memory; } - } - if (!y_data) { - auto y_dims = s_.y_dims; - if (x_dimensions == 3) { - y_dims.erase(y_dims.begin() + 2); + // The following block will be executed in case there has been no change in the shapes of the + // input and the filter compared to the previous run + if (!y_data) { + auto y_dims = s_.y_dims; + if (x_dimensions == 3) { + y_dims.erase(y_dims.begin() + 2); + } + Tensor* Y = context->Output(0, TensorShape(y_dims)); + y_data = reinterpret_cast(Y->template MutableData()); + + // Bail out early if one of the output dimensions is zero. + if (Y->Shape().Size() == 0) { + return Status::OK(); + } } - Tensor* Y = context->Output(0, TensorShape(y_dims)); - y_data = reinterpret_cast(Y->template MutableData()); - } - const auto alpha = Consts::One; - const auto beta = Consts::Zero; + const auto alpha = Consts::One; + const auto beta = Consts::Zero; - IAllocatorUniquePtr workspace = GetScratchBuffer(s_.workspace_bytes); + IAllocatorUniquePtr workspace = GetScratchBuffer(s_.workspace_bytes); - CUDNN_RETURN_IF_ERROR( - cudnnConvolutionBackwardData( - CudnnHandle(), - &alpha, - s_.filter_desc, - w_data, - s_.x_tensor, - x_data, - s_.conv_desc, - s_.algo, - workspace.get(), - s_.workspace_bytes, - &beta, - s_.y_tensor, - y_data)); + CUDNN_RETURN_IF_ERROR( + cudnnConvolutionBackwardData( + CudnnHandle(), + &alpha, + s_.filter_desc, + w_data, + s_.x_tensor, + x_data, + s_.conv_desc, + s_.algo, + workspace.get(), + s_.workspace_bytes, + &beta, + s_.y_tensor, + y_data)); - if (has_bias) { - const Tensor* B = dynamic_padding ? context->Input(3) : context->Input(2); - auto b_data = reinterpret_cast(B->template Data()); - CUDNN_RETURN_IF_ERROR(cudnnAddTensor(CudnnHandle(), &alpha, s_.b_tensor, b_data, &alpha, s_.y_tensor, y_data)); + if (has_bias) { + const Tensor* B = dynamic_padding ? context->Input(3) : context->Input(2); + auto b_data = reinterpret_cast(B->template Data()); + CUDNN_RETURN_IF_ERROR(cudnnAddTensor(CudnnHandle(), &alpha, s_.b_tensor, b_data, &alpha, s_.y_tensor, y_data)); + } } return Status::OK(); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index de6874b0d1..f22d097683 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -69,7 +69,7 @@ struct OperatorRegistrationInformation gsl::span supportedTensorDataTypes; DmlGraphSupport dmlGraphSupport; - std::vector requiredConstantCpuInputs; + std::pair, int> requiredConstantCpuInputs = {{}, 0}; // For use by operators such as Sum, which may require multiple calls to DML, in which case they // can't be represented as nodes in an optimized graph yet. @@ -238,59 +238,66 @@ DML_OP_EXTERN_QUERY_FUNCTION(MaxPool); DML_OP_EXTERN_QUERY_FUNCTION(Slice); DML_OP_EXTERN_QUERY_FUNCTION(Resize); -const static char* const typeNameListDefault[1] = {"T"}; -const static char* const typeNameListTwo[2] = { "T1", "T2" }; -const static char* const typeNameListThree[3] = { "T1", "T2", "T3" }; -const static char* const typeNameListFour[4] = { "T1", "T2", "T3", "T4" }; -const static char* const typeNameListTopK[2] = { "T", "I" }; -const static char* const typeNameListLogicalComparison[2] = { "T", "T1" }; -const static char* const typeNameListConstantOfShape[2] = { "T1", "T2" }; -const static char* const typeNameListScatterGather[2] = { "T", "Tind" }; -const static char* const typeNameListScatterGatherND[1] = { "T" }; // Tind is curiously missing, only allowing 64-bit. -const static char* const typeNameListSlice10[2] = { "T", "Tind" }; -const static char* const typeNameListWhere[2] = { "B", "T" }; -const static char* const typeNameListEyeLike[1] = { "T2" }; -const static SupportedTensorDataTypes supportedTypeListAll[1] = {SupportedTensorDataTypes::All}; -const static SupportedTensorDataTypes supportedTypeListFloat32[1] = {SupportedTensorDataTypes::Float32}; -const static SupportedTensorDataTypes supportedTypeListFloat16to32[1] = {SupportedTensorDataTypes::Float16to32}; -const static SupportedTensorDataTypes supportedTypeListFloat16to32Int32[1] = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::UInt32}; -const static SupportedTensorDataTypes supportedTypeListInt8to32[1] = {SupportedTensorDataTypes::Int8to32}; -const static SupportedTensorDataTypes supportedTypeListInt32to64AndFloat16to32[1] = {SupportedTensorDataTypes::Int32to64|SupportedTensorDataTypes::Float16to32}; -const static SupportedTensorDataTypes supportedTypeListNumericDefault[1] = { SupportedTensorDataTypes::NumericDefault }; -const static SupportedTensorDataTypes supportedTypeListAllScalars[1] = { SupportedTensorDataTypes::AllScalars }; -const static SupportedTensorDataTypes supportedTypeListBool[1] = {SupportedTensorDataTypes::Bool}; -const static SupportedTensorDataTypes supportedTypeListTopK[2] = {SupportedTensorDataTypes::NumericDefault, SupportedTensorDataTypes::Int64}; -const static SupportedTensorDataTypes supportedTypeListIndices[1] = { SupportedTensorDataTypes::Int32|SupportedTensorDataTypes::Int64 }; -const static SupportedTensorDataTypes supportedTypeListCast[2] = { SupportedTensorDataTypes::AllScalars, SupportedTensorDataTypes::AllScalars }; -const static SupportedTensorDataTypes supportedTypeListScalars8to32[1] = { SupportedTensorDataTypes::Scalars8to32 }; -const static SupportedTensorDataTypes supportedTypeListScatterGather[2] = { SupportedTensorDataTypes::Scalars8to32, SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::Int64 }; -const static SupportedTensorDataTypes supportedTypeListScatterGatherND[1] = { SupportedTensorDataTypes::Scalars8to32 }; -const static SupportedTensorDataTypes supportedTypeListSlice10[2] = { SupportedTensorDataTypes::NumericDefault, SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::Int64 }; -const static SupportedTensorDataTypes supportedTypeListQuantizeLinear[2] = { SupportedTensorDataTypes::Float32 | SupportedTensorDataTypes::Int32, SupportedTensorDataTypes::UInt8 | SupportedTensorDataTypes::Int8 }; -const static SupportedTensorDataTypes supportedTypeListDequantizeLinear[2] = { SupportedTensorDataTypes::Float32, SupportedTensorDataTypes::UInt8 | SupportedTensorDataTypes::Int8 | SupportedTensorDataTypes::Int32 }; -const static SupportedTensorDataTypes supportedTypeListQuantize[2] = { SupportedTensorDataTypes::Float32, SupportedTensorDataTypes::UInt8 }; -const static SupportedTensorDataTypes supportedTypeListIsNan[2] = { SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::Bool }; -const static SupportedTensorDataTypes supportedTypeListIsInf[2] = { SupportedTensorDataTypes::Float32, SupportedTensorDataTypes::Bool }; -const static SupportedTensorDataTypes supportedTypeListConstantOfShape[2] = { SupportedTensorDataTypes::Int32|SupportedTensorDataTypes::Int64, SupportedTensorDataTypes::Float16to32 }; -const static SupportedTensorDataTypes supportedTypeListWhere[2] = { SupportedTensorDataTypes::Bool, SupportedTensorDataTypes::AllScalars }; -const static SupportedTensorDataTypes supportedTypeListOneHot[3] = /* indices, depth, values */ { SupportedTensorDataTypes::Int32to64, SupportedTensorDataTypes::AllScalars, SupportedTensorDataTypes::Scalars8to32 }; -const static SupportedTensorDataTypes supportedTypeListLogicalComparison7[2] = /* A&B,C */ { SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::Bool }; -const static SupportedTensorDataTypes supportedTypeListLogicalComparison9[2] = /* A&B,C */ { SupportedTensorDataTypes::NumericDefault, SupportedTensorDataTypes::Bool }; -const static SupportedTensorDataTypes supportedTypeListSigned[1] = { SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::Int16 | SupportedTensorDataTypes::Int8 }; -const static SupportedTensorDataTypes supportedTypeListRange[1] = {SupportedTensorDataTypes::Int16|SupportedTensorDataTypes::Int32|SupportedTensorDataTypes::Float32}; -const static SupportedTensorDataTypes supportedTypeListInteger[3] = {SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int32 }; -const static SupportedTensorDataTypes supportedTypeListQLinearMatMul[3] = { +constexpr static std::array typeNameListDefault = {"T"}; +constexpr static std::array typeNameListTwo = { "T1", "T2" }; +constexpr static std::array typeNameListThree = { "T1", "T2", "T3" }; +constexpr static std::array typeNameListFour = { "T1", "T2", "T3", "T4" }; +constexpr static std::array typeNameListTopK = { "T", "I" }; +constexpr static std::array typeNameListLogicalComparison = { "T", "T1" }; +constexpr static std::array typeNameListConstantOfShape = { "T1", "T2" }; +constexpr static std::array typeNameListScatterGather = { "T", "Tind" }; +constexpr static std::array typeNameListScatterGatherND = { "T" }; // Tind is curiously missing, only allowing 64-bit. +constexpr static std::array typeNameListSlice10 = { "T", "Tind" }; +constexpr static std::array typeNameListWhere = { "B", "T" }; +constexpr static std::array typeNameListEyeLike = { "T2" }; +constexpr static std::array supportedTypeListAll = {SupportedTensorDataTypes::All}; +constexpr static std::array supportedTypeListFloat32 = {SupportedTensorDataTypes::Float32}; +constexpr static std::array supportedTypeListFloat16to32 = {SupportedTensorDataTypes::Float16to32}; +constexpr static std::array supportedTypeListFloat16to32Int32 = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::UInt32}; +constexpr static std::array supportedTypeListInt8to32 = {SupportedTensorDataTypes::Int8to32}; +constexpr static std::array supportedTypeListInt32to64AndFloat16to32 = {SupportedTensorDataTypes::Int32to64|SupportedTensorDataTypes::Float16to32}; +constexpr static std::array supportedTypeListNumericDefault = { SupportedTensorDataTypes::NumericDefault }; +constexpr static std::array supportedTypeListAllScalars = { SupportedTensorDataTypes::AllScalars }; +constexpr static std::array supportedTypeListBool = {SupportedTensorDataTypes::Bool}; +constexpr static std::array supportedTypeListTopK = {SupportedTensorDataTypes::NumericDefault, SupportedTensorDataTypes::Int64}; +constexpr static std::array supportedTypeListIndices = { SupportedTensorDataTypes::Int32|SupportedTensorDataTypes::Int64 }; +constexpr static std::array supportedTypeListCast = { SupportedTensorDataTypes::AllScalars, SupportedTensorDataTypes::AllScalars }; +constexpr static std::array supportedTypeListScalars8to32 = { SupportedTensorDataTypes::Scalars8to32 }; +constexpr static std::array supportedTypeListScatterGather = { SupportedTensorDataTypes::Scalars8to32, SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::Int64 }; +constexpr static std::array supportedTypeListScatterGatherND = { SupportedTensorDataTypes::Scalars8to32 }; +constexpr static std::array supportedTypeListSlice10 = { SupportedTensorDataTypes::NumericDefault, SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::Int64 }; +constexpr static std::array supportedTypeListQuantizeLinear = { SupportedTensorDataTypes::Float32 | SupportedTensorDataTypes::Int32, SupportedTensorDataTypes::UInt8 | SupportedTensorDataTypes::Int8 }; +constexpr static std::array supportedTypeListDequantizeLinear = { SupportedTensorDataTypes::Float32, SupportedTensorDataTypes::UInt8 | SupportedTensorDataTypes::Int8 | SupportedTensorDataTypes::Int32 }; +constexpr static std::array supportedTypeListQuantize = { SupportedTensorDataTypes::Float32, SupportedTensorDataTypes::UInt8 }; +constexpr static std::array supportedTypeListIsNan = { SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::Bool }; +constexpr static std::array supportedTypeListIsInf = { SupportedTensorDataTypes::Float32, SupportedTensorDataTypes::Bool }; +constexpr static std::array supportedTypeListConstantOfShape = { SupportedTensorDataTypes::Int32|SupportedTensorDataTypes::Int64, SupportedTensorDataTypes::Float16to32 }; +constexpr static std::array supportedTypeListWhere = { SupportedTensorDataTypes::Bool, SupportedTensorDataTypes::AllScalars }; +constexpr static std::array supportedTypeListOneHot = /* indices, depth, values */ { SupportedTensorDataTypes::Int32to64, SupportedTensorDataTypes::AllScalars, SupportedTensorDataTypes::Scalars8to32 }; +constexpr static std::array supportedTypeListLogicalComparison7 = /* A&B,C */ { SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::Bool }; +constexpr static std::array supportedTypeListLogicalComparison9 = /* A&B,C */ { SupportedTensorDataTypes::NumericDefault, SupportedTensorDataTypes::Bool }; +constexpr static std::array supportedTypeListSigned = { SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::Int16 | SupportedTensorDataTypes::Int8 }; +constexpr static std::array supportedTypeListRange = {SupportedTensorDataTypes::Int16|SupportedTensorDataTypes::Int32|SupportedTensorDataTypes::Float32}; +constexpr static std::array supportedTypeListInteger = {SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int32 }; +constexpr static std::array supportedTypeListQLinearMatMul = { SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8 }; -const static SupportedTensorDataTypes supportedTypeListQLinearConv[4] = { +constexpr static std::array supportedTypeListQLinearConv = { SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int8|SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Int32 }; +template +constexpr auto requiredConstantCpuInputs(Args... args) +{ + std::array inputs = {static_cast(args)...}; + return std::make_pair(inputs, static_cast(sizeof...(args))); +} + // Define a single row of registration information. #define REG_INFO(version, operatorName, ...) \ #operatorName, OnnxOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kOnnxDomain, Create##operatorName, ShapeInferenceFunction, false, false, ##__VA_ARGS__, @@ -314,7 +321,7 @@ const static SupportedTensorDataTypes supportedTypeListQLinearConv[4] = { #define REG_INFO_MSDML(version, operatorName, ...) \ #operatorName, MsftOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kMSDmlDomain, Create##operatorName, ShapeInferenceFunction, false, false, ##__VA_ARGS__, -const static OperatorRegistrationInformation operatorRegistrationInformationTable[] = +constexpr static OperatorRegistrationInformation operatorRegistrationInformationTable[] = { /// Domain/Type, Ver, Name, TypeNames, Types, Graph Support, Required const CPU inputs, /// Input count required for graph support, @@ -330,9 +337,9 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 11, AveragePool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, GlobalAveragePool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, - {REG_INFO( 8, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {}, std::nullopt, QueryMaxPool)}, - {REG_INFO( 10, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {}, std::nullopt, QueryMaxPool)}, - {REG_INFO( 11, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {}, std::nullopt, QueryMaxPool)}, + {REG_INFO( 8, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(), std::nullopt, QueryMaxPool)}, + {REG_INFO( 10, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(), std::nullopt, QueryMaxPool)}, + {REG_INFO( 11, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(), std::nullopt, QueryMaxPool)}, {REG_INFO( 7, GlobalMaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, LpPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -349,7 +356,7 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 7, RNN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, {REG_INFO( 7, GRU, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, {REG_INFO( 7, LSTM, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, - {REG_INFO_MS( 1, ConvTransposeWithDynamicPads, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {2})}, + {REG_INFO_MS( 1, ConvTransposeWithDynamicPads, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(2))}, // Data Reorganization Layers {REG_INFO( 7, Split, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, @@ -358,16 +365,16 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 7, Concat, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO( 11, Concat, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, // Adds negative axis. {REG_INFO_VER( 7, Slice, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, - {REG_INFO_VER( 10, Slice, typeNameListSlice10, supportedTypeListSlice10, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {1, 2, 3, 4}, std::nullopt, QuerySlice)}, // Adds negative axes. - {REG_INFO_VER( 11, Slice, typeNameListSlice10, supportedTypeListSlice10, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {1, 2, 3, 4}, std::nullopt, QuerySlice)}, + {REG_INFO_VER( 10, Slice, typeNameListSlice10, supportedTypeListSlice10, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(1, 2, 3, 4), std::nullopt, QuerySlice)}, // Adds negative axes. + {REG_INFO_VER( 11, Slice, typeNameListSlice10, supportedTypeListSlice10, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(1, 2, 3, 4), std::nullopt, QuerySlice)}, {REG_INFO_VER( 7, Pad, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, - {REG_INFO_VER( 11, Pad, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {1, 2} /*pads, value*/)}, // https://microsoft.visualstudio.com/OS/_workitems/edit/26007728 + {REG_INFO_VER( 11, Pad, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(1, 2) /*pads, value*/)}, // https://microsoft.visualstudio.com/OS/_workitems/edit/26007728 {REG_INFO( 7, SpaceToDepth, typeNameListDefault, supportedTypeListScalars8to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO( 7, DepthToSpace, typeNameListDefault, supportedTypeListScalars8to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO( 11, DepthToSpace, typeNameListDefault, supportedTypeListScalars8to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, - {REG_INFO( 7, Tile, typeNameListDefault, supportedTypeListScalars8to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {1})}, - {REG_INFO( 8, Expand, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {1})}, - {REG_INFO( 9, ConstantOfShape, typeNameListConstantOfShape, supportedTypeListConstantOfShape, DmlGraphSupport::NotSupported, {0})}, + {REG_INFO( 7, Tile, typeNameListDefault, supportedTypeListScalars8to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(1))}, + {REG_INFO( 8, Expand, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(1))}, + {REG_INFO( 9, ConstantOfShape, typeNameListConstantOfShape, supportedTypeListConstantOfShape, DmlGraphSupport::NotSupported, requiredConstantCpuInputs(0))}, {REG_INFO( 7, Gather, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, {REG_INFO( 11, Gather, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, {REG_INFO( 11, GatherElements, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, @@ -387,7 +394,7 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO_ID( 11, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO_ID( 7, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO_ID( 11, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, - {REG_INFO_ID( 7, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {1})}, + {REG_INFO_ID( 7, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(1))}, // Elementwise {REG_INFO( 7, Sqrt, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -399,19 +406,19 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 7, Ceil, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Floor, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO_VER( 7, Clip, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO_VER( 11, Clip, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {1,2})}, + {REG_INFO_VER( 11, Clip, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1,2))}, {REG_INFO( 7, Add, typeNameListDefault, supportedTypeListFloat16to32Int32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Sub, typeNameListDefault, supportedTypeListFloat16to32Int32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Mul, typeNameListDefault, supportedTypeListFloat16to32Int32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Div, typeNameListDefault, supportedTypeListFloat16to32Int32, DmlGraphSupport::Supported)}, - {REG_INFO( 7, Sum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, - {REG_INFO( 8, Sum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, - {REG_INFO( 7, Mean, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, - {REG_INFO( 8, Mean, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, - {REG_INFO( 7, Max, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, - {REG_INFO( 8, Max, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, - {REG_INFO( 7, Min, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, - {REG_INFO( 8, Min, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, + {REG_INFO( 7, Sum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, + {REG_INFO( 8, Sum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, + {REG_INFO( 7, Mean, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, + {REG_INFO( 8, Mean, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, + {REG_INFO( 7, Max, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, + {REG_INFO( 8, Max, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, + {REG_INFO( 7, Min, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, + {REG_INFO( 8, Min, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, {REG_INFO( 7, Cos, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Sin, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Tan, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -475,10 +482,10 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 7, Crop, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, ImageScaler, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO_VER( 7, Upsample, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO_VER( 9, Upsample, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {1} /*scales*/)}, - {REG_INFO_VER( 10, Upsample, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {1} /*scales*/)}, - {REG_INFO_VER( 10, Resize, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {1} /*scales*/)}, - {REG_INFO_VER( 11, Resize, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {1, 2, 3} /*roi, scales, sizes*/, std::nullopt, QueryResize)}, + {REG_INFO_VER( 9, Upsample, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1) /*scales*/)}, + {REG_INFO_VER( 10, Upsample, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1) /*scales*/)}, + {REG_INFO_VER( 10, Resize, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1) /*scales*/)}, + {REG_INFO_VER( 11, Resize, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3) /*roi, scales, sizes*/, std::nullopt, QueryResize)}, // Activation Functions {REG_INFO( 7, Sigmoid, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -513,10 +520,10 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 7, MemcpyFromHost, typeNameListDefault, supportedTypeListAll)}, {REG_INFO( 7, MemcpyToHost, typeNameListDefault, supportedTypeListAll)}, {REG_INFO_VER( 7, TopK, typeNameListTopK, supportedTypeListTopK, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, - {REG_INFO_VER( 10, TopK, typeNameListTopK, supportedTypeListTopK, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {1})}, - {REG_INFO_VER( 11, TopK, typeNameListTopK, supportedTypeListTopK, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, {1})}, - {REG_INFO( 9, OneHot, typeNameListThree, supportedTypeListOneHot, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp, {1})}, - {REG_INFO( 11, OneHot, typeNameListThree, supportedTypeListOneHot, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp, {1})}, + {REG_INFO_VER( 10, TopK, typeNameListTopK, supportedTypeListTopK, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(1))}, + {REG_INFO_VER( 11, TopK, typeNameListTopK, supportedTypeListTopK, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides, requiredConstantCpuInputs(1))}, + {REG_INFO( 9, OneHot, typeNameListThree, supportedTypeListOneHot, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp, requiredConstantCpuInputs(1))}, + {REG_INFO( 11, OneHot, typeNameListThree, supportedTypeListOneHot, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp, requiredConstantCpuInputs(1))}, // Fused operators {REG_INFO_MSDML(1, FusedConv, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -527,18 +534,18 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO_MSDML(1, FusedGemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO_MSDML(1, FusedMatMul, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO_MSDML(1, FusedAdd, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO_MSDML(1, FusedSum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {}, 2)}, + {REG_INFO_MSDML(1, FusedSum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, {REG_INFO( 10, IsInf, typeNameListTwo, supportedTypeListIsInf, DmlGraphSupport::Supported)}, {REG_INFO( 10, Mod, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported)}, {REG_INFO( 11, BitShift, typeNameListDefault, supportedTypeListInt8to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, Round, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 10, ReverseSequence, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, - {REG_INFO( 11, CumSum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, {1})}, - {REG_INFO( 11, Range, typeNameListDefault, supportedTypeListRange, DmlGraphSupport::Supported, {0,1,2})}, + {REG_INFO( 11, CumSum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + {REG_INFO( 11, Range, typeNameListDefault, supportedTypeListRange, DmlGraphSupport::Supported, requiredConstantCpuInputs(0,1,2))}, - {REG_INFO( 9, MaxUnpool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp, {2})}, - {REG_INFO( 11, MaxUnpool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp, {2})}, // 11 is identical to 9. + {REG_INFO( 9, MaxUnpool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp, requiredConstantCpuInputs(2))}, + {REG_INFO( 11, MaxUnpool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp, requiredConstantCpuInputs(2))}, // 11 is identical to 9. {REG_INFO( 10, QLinearConv, typeNameListFour, supportedTypeListQLinearConv, DmlGraphSupport::NotSupported)}, {REG_INFO( 10, QLinearMatMul, typeNameListThree, supportedTypeListQLinearMatMul, DmlGraphSupport::NotSupported)}, @@ -652,8 +659,8 @@ void RegisterDmlOperators(IMLOperatorRegistry* registry) supportedWith64BitTensorsVia32BitStrides, supportedWith64BitTensorsVia32BitStridesFromAnyEp, prefer64BitTensorsDirectly, - information.requiredConstantCpuInputs.empty() ? nullptr : information.requiredConstantCpuInputs.data(), - static_cast(information.requiredConstantCpuInputs.size()) + information.requiredConstantCpuInputs.first.data(), + static_cast(information.requiredConstantCpuInputs.second) )); } } diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/model_builder.cc b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/model_builder.cc index 6d840eab7a..98844f9c73 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/model_builder.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/model_builder.cc @@ -62,7 +62,7 @@ bool IsValidSupportedNodesVec(const std::vector& supported_node_vec, std::vector> ModelBuilder::GetSupportedNodes() { std::vector> supported_node_vecs; - int32_t android_sdk_ver = nnapi_ ? nnapi_->android_sdk_version : 0; + int32_t android_sdk_ver = GetAndroidSdkVer(); #ifdef __ANDROID__ if (android_sdk_ver < 27) { LOGS_DEFAULT(VERBOSE) << "Android API level " @@ -126,6 +126,7 @@ void ModelBuilder::AddInitializerToSkip(const std::string& tensor_name) { void ModelBuilder::Prepare() { nnapi_model_ = std::unique_ptr(new Model()); THROW_ON_ERROR(nnapi_->ANeuralNetworksModel_create(&nnapi_model_->model_)); + GetTargetDevices(); PreprocessInitializers(); RegisterInitializers(); RegisterModelInputs(); @@ -144,9 +145,40 @@ static size_t GetPaddedByteSize(size_t size) { return (size + kDefaultByteAlignmentForNNAPI - 1) & ~(kDefaultByteAlignmentForNNAPI - 1); } +void ModelBuilder::GetTargetDevices() { + // GetTargetDevices is only supported on API 29+ + if (GetAndroidSdkVer() < 29) + return; + + if (target_device_option_ == TargetDeviceOption::ALL_DEVICES) + return; + + const std::string nnapi_cpu("nnapi-reference"); + uint32_t num_devices = 0; + THROW_ON_ERROR_WITH_NOTE(nnapi_->ANeuralNetworks_getDeviceCount(&num_devices), + "Getting list of available devices"); + + for (uint32_t i = 0; i < num_devices; i++) { + ANeuralNetworksDevice* device = nullptr; + const char* device_name = nullptr; + THROW_ON_ERROR_WITH_NOTE(nnapi_->ANeuralNetworks_getDevice(i, &device), + "Getting list of available devices"); + + THROW_ON_ERROR_WITH_NOTE(nnapi_->ANeuralNetworksDevice_getName(device, &device_name), + "Getting list of available devices"); + + bool device_is_cpu = nnapi_cpu == device_name; + if ((target_device_option_ == TargetDeviceOption::CPU_DISABLED && !device_is_cpu) || + (target_device_option_ == TargetDeviceOption::CPU_ONLY && device_is_cpu)) { + nnapi_target_devices_.push_back(device); + LOGS_DEFAULT(VERBOSE) << "Target device [" << device_name << "] added"; + } + } +} + void ModelBuilder::GetAllInitializers() { for (const auto& tensor : model_proto_.graph().initializer()) { - initializers_.insert({tensor.name(), tensor}); + initializers_.emplace(tensor.name(), tensor); } } @@ -190,7 +222,7 @@ void ModelBuilder::RegisterInitializers() { OperandType operand_type(type, shape); shaper_.AddShape(name, operand_type.dimensions); - auto index = AddNewOperand(name, operand_type); + auto index = AddNewOperand(name, operand_type, false /* is_nhwc */); const size_t size = operand_type.GetOperandBlobByteSize(); const size_t padded_size = GetPaddedByteSize(size); sizeAll += padded_size; @@ -264,7 +296,7 @@ void ModelBuilder::RegisterModelInputs() { OperandType operand_type(type, shape); shaper_.AddShape(input_name, operand_type.dimensions); - auto index = AddNewOperand(input_name, operand_type); + auto index = AddNewOperand(input_name, operand_type, false /* is_nhwc */); input_index_vec_.push_back(index); nnapi_model_->AddInput(input_name, operand_type); @@ -279,8 +311,15 @@ void ModelBuilder::RegisterModelOutputs() { ORT_THROW("The output of graph is not registered" + output_name); } - output_index_vec_.push_back(operand_indices_[output_name]); - nnapi_model_->AddOutput(output_name, operand_types_.at(output_name)); + std::string nnapi_output_name = output_name; + if (IsOperandNHWC(output_name)) { + // We need to transpose the output still in nhwc back to nchw + nnapi_output_name = GetUniqueName(output_name + "_nhwc_to_nchw"); + TransposeNHWCToNCHW(*this, output_name, nnapi_output_name); + } + + output_index_vec_.push_back(operand_indices_[nnapi_output_name]); + nnapi_model_->AddOutput(output_name, nnapi_output_name, operand_types_.at(nnapi_output_name)); } } @@ -290,11 +329,12 @@ void ModelBuilder::RegisterModelShaper() { } uint32_t ModelBuilder::AddNewOperand(const std::string& name, - const android::nn::wrapper::OperandType& operand_type) { + const OperandType& operand_type, + bool is_nhwc) { THROW_ON_ERROR(nnapi_->ANeuralNetworksModel_addOperand( nnapi_model_->model_, &operand_type.operandType)); auto idx = next_index_++; - RegisterOperand(name, idx, operand_type); + RegisterOperand(name, idx, operand_type, is_nhwc); return idx; } @@ -304,12 +344,14 @@ uint32_t ModelBuilder::AddNewNNAPIOperand(const OperandType& operand_type) { return next_index_++; } -void ModelBuilder::RegisterOperand(const std::string& name, - uint32_t index, - const OperandType& operand_type) { +void ModelBuilder::RegisterOperand(const std::string& name, uint32_t index, + const OperandType& operand_type, bool is_nhwc) { operand_indices_[name] = index; - operand_types_.insert({name, operand_type}); + operand_types_.emplace(name, operand_type); operands_.insert(name); + + if (is_nhwc) + RegisterNHWCOperand(name); } void ModelBuilder::SetOperandValue(uint32_t index, @@ -334,7 +376,7 @@ uint32_t ModelBuilder::AddOperandFromPersistMemoryBuffer( const std::string& name, const void* buffer, const android::nn::wrapper::OperandType& operand_type) { shaper_.AddShape(name, operand_type.dimensions); - auto index = AddNewOperand(name, operand_type); + auto index = AddNewOperand(name, operand_type, false /* is_nhwc */); const size_t size = operand_type.GetOperandBlobByteSize(); // for small size operand, the value will be copied @@ -369,10 +411,11 @@ void ModelBuilder::AddOperations() { void ModelBuilder::AddOperation(int op, const std::vector& input_indices, const std::vector& output_names, - const std::vector& types) { + const std::vector& types, + const std::vector& is_nhwc_vec) { std::vector output_indices; for (size_t i = 0; i < types.size(); i++) { - output_indices.push_back(AddNewOperand(output_names[i], types[i])); + output_indices.push_back(AddNewOperand(output_names[i], types[i], is_nhwc_vec[i])); } THROW_ON_ERROR_WITH_NOTE( @@ -393,7 +436,8 @@ std::unique_ptr ModelBuilder::Compile() { &output_index_vec_[0]), "on identifyInputsAndOutputs"); - if (use_fp16_) { + // relax fp32tofp16 is only available on API 28+ + if (use_fp16_ && GetAndroidSdkVer() > 27) { THROW_ON_ERROR_WITH_NOTE( nnapi_->ANeuralNetworksModel_relaxComputationFloat32toFloat16( nnapi_model_->model_, true), @@ -404,9 +448,17 @@ std::unique_ptr ModelBuilder::Compile() { nnapi_->ANeuralNetworksModel_finish(nnapi_model_->model_), "on model finish"); - THROW_ON_ERROR_WITH_NOTE( - nnapi_->ANeuralNetworksCompilation_create(nnapi_model_->model_, &nnapi_model_->compilation_), - "on create"); + if (!nnapi_target_devices_.empty()) { + THROW_ON_ERROR_WITH_NOTE( + nnapi_->ANeuralNetworksCompilation_createForDevices( + nnapi_model_->model_, nnapi_target_devices_.data(), + nnapi_target_devices_.size(), &nnapi_model_->compilation_), + "on createForDevices"); + } else { + THROW_ON_ERROR_WITH_NOTE( + nnapi_->ANeuralNetworksCompilation_create(nnapi_model_->model_, &nnapi_model_->compilation_), + "on create"); + } THROW_ON_ERROR_WITH_NOTE( nnapi_->ANeuralNetworksCompilation_setPreference( @@ -475,5 +527,41 @@ std::string ModelBuilder::GetUniqueName(const std::string& base_name) { return unique_name; } +void ModelBuilder::RegisterNHWCOperand(const std::string& name) { + nhwc_operands_.insert(name); +} + +bool ModelBuilder::IsOperandNHWC(const std::string& name) { + return Contains(nhwc_operands_, name); +} + +bool ModelBuilder::GetNCHWOperand(const std::string& nhwc_name, std::string& nchw_name) { + if (Contains(nhwc_to_nchw_map_, nhwc_name)) { + nchw_name = nhwc_to_nchw_map_[nhwc_name]; + return true; + } + return false; +} + +bool ModelBuilder::GetNHWCOperand(const std::string& nchw_name, std::string& nhwc_name) { + if (Contains(nchw_to_nhwc_map_, nchw_name)) { + nhwc_name = nchw_to_nhwc_map_[nchw_name]; + return true; + } + return false; +} + +void ModelBuilder::SetNHWCToNCHWOperandMap(const std::string& nhwc_name, + const std::string& nchw_name) { + ORT_ENFORCE(!Contains(nhwc_to_nchw_map_, nhwc_name), "A previous nchw to nhwc map exists"); + nhwc_to_nchw_map_[nhwc_name] = nchw_name; +} + +void ModelBuilder::SetNCHWToNHWCOperandMap(const std::string& nchw_name, + const std::string& nhwc_name) { + ORT_ENFORCE(!Contains(nchw_to_nhwc_map_, nchw_name), "A previous nchw to nhwc map exists"); + nchw_to_nhwc_map_[nchw_name] = nhwc_name; +} + } // namespace nnapi } // namespace onnxruntime \ No newline at end of file diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/model_builder.h b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/model_builder.h index a475ab6659..d9ca4a1b69 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/model_builder.h +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/model_builder.h @@ -17,6 +17,15 @@ class ModelBuilder { public: using Shape = Shaper::Shape; + enum class TargetDeviceOption : int8_t { + ALL_DEVICES, // use all avaliable target devices + /* TODO support this option + SINGLE_DEVICE, // use a single target device, must be given + */ + CPU_DISABLED, // use all avaliable target devices except CPU + CPU_ONLY, // use CPU only + }; + ModelBuilder(ONNX_NAMESPACE::ModelProto& model_proto); ~ModelBuilder() = default; @@ -29,7 +38,8 @@ class ModelBuilder { // Add an NNAPI operation (operator) void AddOperation(int op, const std::vector& input_indices, const std::vector& output_names, - const std::vector& types); + const std::vector& types, + const std::vector& is_nhwc_vec); // Find if an output has a fuseable activation (Relu) int32_t FindActivation(const std::string& output); @@ -48,7 +58,8 @@ class ModelBuilder { // Register informations for a particular operand void RegisterOperand(const std::string& name, uint32_t index, - const android::nn::wrapper::OperandType& operand_type); + const android::nn::wrapper::OperandType& operand_type, + bool is_nhwc); // Generate an unique name for intermediate result std::string GetUniqueName(const std::string& base_name); @@ -84,6 +95,18 @@ class ModelBuilder { const ONNX_NAMESPACE::ModelProto& GetOnnxModel() const { return model_proto_; } + void RegisterNHWCOperand(const std::string& name); + bool IsOperandNHWC(const std::string& name); + + // Get the operand transposed to nchw/nhwc from given nhwc/nchw operand, if it exists + bool GetNCHWOperand(const std::string& nhwc_name, std::string& nchw_name); + bool GetNHWCOperand(const std::string& nchw_name, std::string& nhwc_name); + + void SetNHWCToNCHWOperandMap(const std::string& nhwc_name, + const std::string& nchw_name); + void SetNCHWToNHWCOperandMap(const std::string& nchw_name, + const std::string& nhwc_name); + private: const NnApi* nnapi_{nullptr}; ONNX_NAMESPACE::ModelProto& model_proto_; @@ -91,7 +114,7 @@ class ModelBuilder { uint32_t name_token_{0}; - bool use_nchw_{true}; + bool use_nchw_{false}; bool use_fp16_{false}; android::nn::wrapper::ExecutePreference exe_pref_{ android::nn::wrapper::ExecutePreference::PREFER_FAST_SINGLE_ANSWER}; @@ -109,11 +132,21 @@ class ModelBuilder { std::unordered_map> op_builders_; + // Operands in nhwc + std::unordered_set nhwc_operands_; + + // Maps between nhwc and nchw, and vice versa + std::unordered_map nhwc_to_nchw_map_; + std::unordered_map nchw_to_nhwc_map_; + std::vector input_index_vec_; std::vector output_index_vec_; std::unordered_set unique_names_; + TargetDeviceOption target_device_option_{TargetDeviceOption::ALL_DEVICES}; + std::vector nnapi_target_devices_; + uint32_t next_index_ = 0; bool IsNodeSupported(const ONNX_NAMESPACE::NodeProto& node); @@ -121,6 +154,7 @@ class ModelBuilder { // Convert the onnx model to ANeuralNetworksModel void Prepare(); + void GetTargetDevices(); void GetAllInitializers(); void PreprocessInitializers(); void RegisterInitializers(); @@ -134,7 +168,8 @@ class ModelBuilder { uint32_t AddNewNNAPIOperand(const android::nn::wrapper::OperandType& type); uint32_t AddNewOperand(const std::string& name, - const android::nn::wrapper::OperandType& operand_type); + const android::nn::wrapper::OperandType& operand_type, + bool is_nhwc); IOpBuilder* GetOpBuilder(const ONNX_NAMESPACE::NodeProto& node); }; diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/op_builder.cc b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/op_builder.cc index 4a3c14dd9a..3b69ee44a6 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/op_builder.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/op_builder.cc @@ -30,12 +30,87 @@ const float* GetTensorFloatData(const ONNX_NAMESPACE::TensorProto& tensor) { : tensor.float_data().data(); } +void AddTransposeOperator(ModelBuilder& model_builder, + const std::string& input, + const std::string& perm_name, + vector perm, + const std::string& output, + bool output_is_nhwc) { + auto& shaper(model_builder.GetShaper()); + const auto& operand_indices(model_builder.GetOperandIndices()); + const auto& operand_types(model_builder.GetOperandTypes()); + + std::vector input_indices; + input_indices.push_back(operand_indices.at(input)); // input + + ModelBuilder::Shape perm_dimen = {SafeInt(perm.size())}; + OperandType perm_operand_type(Type::TENSOR_INT32, perm_dimen); + uint32_t perm_idx = model_builder.AddOperandFromPersistMemoryBuffer( + perm_name, perm.data(), perm_operand_type); + + input_indices.push_back(perm_idx); // permutation + shaper.Transpose(input, perm, output); + const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); + model_builder.AddOperation(ANEURALNETWORKS_TRANSPOSE, input_indices, {output}, + {output_operand_type}, {output_is_nhwc}); +} + +void TransposeBetweenNCHWAndNHWC(ModelBuilder& model_builder, + const std::string& input, + const std::string& output, + bool nchw_to_nhwc) { + ORT_ENFORCE(!model_builder.UseNCHW(), "model_builder.UseNCHW() is on"); + const auto& shaper(model_builder.GetShaper()); + ORT_ENFORCE( + 4 == shaper[input].size(), + "TransposeNCHWToNHWC input has to be a 4d tensor, actual dimensions: " + + std::to_string(shaper[input].size())); + + std::string perm_name; + vector perm; + if (nchw_to_nhwc) { + perm_name = model_builder.GetUniqueName(input + "nchw_to_nhwc_perm"); + perm = {0, 2, 3, 1}; + } else { // nhwc_to_nchw + perm_name = model_builder.GetUniqueName(input + "nhwc_to_nchw_perm"); + perm = {0, 3, 1, 2}; + } + + AddTransposeOperator(model_builder, input, perm_name, perm, output, nchw_to_nhwc); + + if (nchw_to_nhwc) { + model_builder.SetNCHWToNHWCOperandMap(input, output); + } else { // nhwc_to_nchw + model_builder.SetNHWCToNCHWOperandMap(input, output); + } + + LOGS_DEFAULT(VERBOSE) << "Operand [" << input << "] with shape " + << Shape2String(shaper[input]) + << " is transposed " + << (nchw_to_nhwc ? "nchw_to_nhwc" : "nhwc_to_nchw") + << " to [" << output << "] with shape " + << Shape2String(shaper[output]); +} + +void TransposeNHWCToNCHW(ModelBuilder& model_builder, + const std::string& input, + const std::string& output) { + TransposeBetweenNCHWAndNHWC(model_builder, input, output, false /* nchw_to_nhwc */); +} + +void TransposeNCHWToNHWC(ModelBuilder& model_builder, + const std::string& input, + const std::string& output) { + TransposeBetweenNCHWAndNHWC(model_builder, input, output, true /* nchw_to_nhwc */); +} + void AddBinaryOperator(int32_t op_type, ModelBuilder& model_builder, const std::string& input1, const std::string& input2, int32_t fuse_code, - const std::string& output) { + const std::string& output, + bool output_is_nhwc) { auto& shaper(model_builder.GetShaper()); const auto& operand_indices(model_builder.GetOperandIndices()); const auto& operand_types(model_builder.GetOperandTypes()); @@ -46,41 +121,7 @@ void AddBinaryOperator(int32_t op_type, input_indices.push_back(model_builder.AddOperandFromScalar(fuse_code)); shaper.Eltwise(input1, input2, output); const OperandType output_operand_type(operand_types.at(input1).type, shaper[output]); - model_builder.AddOperation(op_type, input_indices, {output}, {output_operand_type}); -} - -void AddPoolOperator(int32_t op_type, - ModelBuilder& model_builder, - const std::string& input, - const vector& onnx_pads, - const vector& onnx_strides, - const vector& kernel_shape, - int32_t fuse_code, - const std::string& output) { - auto& shaper(model_builder.GetShaper()); - const auto& operand_indices(model_builder.GetOperandIndices()); - const auto& operand_types(model_builder.GetOperandTypes()); - bool use_nchw = model_builder.UseNCHW(); - - std::vector input_indices; - input_indices.push_back(operand_indices.at(input)); - input_indices.push_back(model_builder.AddOperandFromScalar(onnx_pads[1])); - input_indices.push_back(model_builder.AddOperandFromScalar(onnx_pads[3])); - input_indices.push_back(model_builder.AddOperandFromScalar(onnx_pads[0])); - input_indices.push_back(model_builder.AddOperandFromScalar(onnx_pads[2])); - input_indices.push_back(model_builder.AddOperandFromScalar(onnx_strides[1])); - input_indices.push_back(model_builder.AddOperandFromScalar(onnx_strides[0])); - input_indices.push_back(model_builder.AddOperandFromScalar(kernel_shape[1])); - input_indices.push_back(model_builder.AddOperandFromScalar(kernel_shape[0])); - input_indices.push_back(model_builder.AddOperandFromScalar(fuse_code)); - input_indices.push_back(model_builder.AddOperandFromScalar(use_nchw)); - - shaper.Pool(input, - onnx_pads, onnx_strides, kernel_shape, - use_nchw, - output); - const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); - model_builder.AddOperation(op_type, input_indices, {output}, {output_operand_type}); + model_builder.AddOperation(op_type, input_indices, {output}, {output_operand_type}, {output_is_nhwc}); } int GetType(const ONNX_NAMESPACE::ModelProto& model_proto, @@ -159,7 +200,7 @@ Shaper::Shape GetShape(const ONNX_NAMESPACE::ModelProto& model_proto, } enum DataLayout { - L_NCHW = 0, + L_0231 = 0, L_1230 = 1, }; @@ -187,8 +228,8 @@ uint32_t AddInitializerInNewLayout(ModelBuilder& model_builder, auto out_t = shape[0], in_t = shape[1], h_t = shape[2], w_t = shape[3]; ModelBuilder::Shape dest_shape; - if (new_layout == L_NCHW) - dest_shape = {out_t, h_t, w_t, in_t}; // L_NCHW + if (new_layout == L_0231) + dest_shape = {out_t, h_t, w_t, in_t}; // L_0231 else dest_shape = {in_t, h_t, w_t, out_t}; // L_1230 for depthwise conv weight @@ -205,7 +246,7 @@ uint32_t AddInitializerInNewLayout(ModelBuilder& model_builder, w; uint32_t nnapi_idx; - if (new_layout == L_NCHW) { // L_NCHW + if (new_layout == L_0231) { // L_0231 nnapi_idx = out * h_t * w_t * in_t + h * w_t * in_t + w * in_t + @@ -389,12 +430,33 @@ void BinaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, else { ORT_THROW("UnaryOpBuilder, unknown op: " + op); } - const auto& input1 = node.input(0); - const auto& input2 = node.input(1); + std::string input1 = node.input(0); + std::string input2 = node.input(1); + bool input1_is_nhwc = model_builder.IsOperandNHWC(input1); + bool input2_is_nhwc = model_builder.IsOperandNHWC(input2); + bool output_is_nhwc = false; + + if (input1_is_nhwc == input2_is_nhwc) { + output_is_nhwc = input1_is_nhwc; + } else if (input1_is_nhwc) { + // need transpsoe input1 back to nchw + const auto& nhwc_input = node.input(0); + if (!model_builder.GetNCHWOperand(nhwc_input, input1)) { + input1 = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw"); + TransposeNHWCToNCHW(model_builder, nhwc_input, input1); + } + } else { // input2_is_nhwc + // need transpsoe input2 back to nchw + const auto& nhwc_input = node.input(1); + if (!model_builder.GetNCHWOperand(nhwc_input, input2)) { + input2 = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw"); + TransposeNHWCToNCHW(model_builder, nhwc_input, input2); + } + } + const auto& output = node.output(0); int32_t fuse_code = model_builder.FindActivation(output); - AddBinaryOperator(op_code, model_builder, - input1, input2, fuse_code, output); + AddBinaryOperator(op_code, model_builder, input1, input2, fuse_code, output, output_is_nhwc); } #pragma endregion @@ -415,16 +477,17 @@ void ReluOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto& input = node.input(0); const auto& output = node.output(0); + bool output_is_nhwc = model_builder.IsOperandNHWC(input); shaper.Identity(input, output); const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); // skip this relu if it is some op's fuse output if (Contains(model_builder.GetFusedActivations(), node.name())) { - model_builder.RegisterOperand(output, operand_indices.at(input), output_operand_type); + model_builder.RegisterOperand(output, operand_indices.at(input), output_operand_type, output_is_nhwc); } else { std::vector input_indices; input_indices.push_back(operand_indices.at(input)); - model_builder.AddOperation(ANEURALNETWORKS_RELU, input_indices, {output}, {output_operand_type}); + model_builder.AddOperation(ANEURALNETWORKS_RELU, input_indices, {output}, {output_operand_type}, {output_is_nhwc}); } } @@ -462,31 +525,35 @@ bool TransposeOpBuilder::IsOpSupportedImpl(ModelBuilder& model_builder, void TransposeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const ONNX_NAMESPACE::NodeProto& node) { auto& shaper(model_builder.GetShaper()); - const auto& operand_indices(model_builder.GetOperandIndices()); - const auto& operand_types(model_builder.GetOperandTypes()); + + auto input = node.input(0); + const auto& output = node.output(0); NodeAttrHelper helper(node); - - const auto& input = node.input(0); - std::vector input_indices; - input_indices.push_back(operand_indices.at(input)); // input - vector perm = helper.Get("perm", vector()); auto input_dims = shaper[input].size(); if (perm.empty()) { for (int32_t i = input_dims - 1; i >= 0; i--) perm.push_back(i); + } else { + ORT_ENFORCE(perm.size() == input_dims, "Perm and input should have same dimension"); } - ModelBuilder::Shape perm_dimen = {SafeInt(input_dims)}; - std::string perm_name = model_builder.GetUniqueName(node.name() + input + "perm"); - OperandType perm_operand_type(Type::TENSOR_INT32, perm_dimen); - uint32_t perm_idx = model_builder.AddOperandFromPersistMemoryBuffer(perm_name, perm.data(), perm_operand_type); - input_indices.push_back(perm_idx); + if (model_builder.IsOperandNHWC(input)) { + ORT_ENFORCE(input_dims == 4, "Only 4D shape can be nhwc"); - const auto& output = node.output(0); - shaper.Transpose(input, perm, output); - const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); - model_builder.AddOperation(ANEURALNETWORKS_TRANSPOSE, input_indices, {output}, {output_operand_type}); + // we are using nhwc here, but the axis is in nchw, need to transpose axis from nchw to nhwc + const int32_t axis_nchw_to_nhwc[4]{0, 3, 1, 2}; + for (size_t i = 0; i < perm.size(); i++) + perm[i] = axis_nchw_to_nhwc[perm[i]]; + } + + std::string perm_name = model_builder.GetUniqueName(node.name() + input + "perm"); + + // It is possible this onnx transpose operator can be nchw->nhwc, but so far I don't see + // any scenario will do this since onnx is nchw only, assume the output is always not nhwc + // even it is, there will be extra transpose in the onnx model to convert it back to nchw + // before conv/pool/... operators + AddTransposeOperator(model_builder, input, perm_name, perm, output, false /* is_nhwc */); } #pragma endregion op_transpose @@ -551,7 +618,17 @@ void ReshapeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto& operand_types(model_builder.GetOperandTypes()); const auto& initializers(model_builder.GetInitializerTensors()); - const auto& input = node.input(0); + auto input = node.input(0); + + if (model_builder.IsOperandNHWC(input)) { + // We want to transpose nhwc operand back to nchw before reshape + const auto& nhwc_input = node.input(0); + if (!model_builder.GetNCHWOperand(nhwc_input, input)) { + input = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw"); + TransposeNHWCToNCHW(model_builder, nhwc_input, input); + } + } + const auto& output = node.output(0); std::vector input_indices; input_indices.push_back(operand_indices.at(input)); // input @@ -576,7 +653,8 @@ void ReshapeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, shaper.Reshape(input, shape, output); const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); - model_builder.AddOperation(ANEURALNETWORKS_RESHAPE, input_indices, {output}, {output_operand_type}); + model_builder.AddOperation(ANEURALNETWORKS_RESHAPE, input_indices, + {output}, {output_operand_type}, {false}); } #pragma endregion op_reshape @@ -676,10 +754,13 @@ void BatchNormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_buil const auto tensor_b_name = model_builder.GetUniqueName(node.name() + input + "_imm_b"); const auto tensor_imm_product_name = model_builder.GetUniqueName(node.name() + input + "_imm_mul"); ModelBuilder::Shape tensor_a_dimen; - if (model_builder.UseNCHW()) - tensor_a_dimen = {size, 1, 1}; // {C, H, W} - else + + bool input_is_nhwc = model_builder.IsOperandNHWC(input); + bool output_is_nhwc = input_is_nhwc; + if (input_is_nhwc) tensor_a_dimen = {size}; + else // input is nchw + tensor_a_dimen = {size, 1, 1}; // {C, H, W} shaper.AddShape(tensor_a_name, tensor_a_dimen); shaper.AddShape(tensor_b_name, tensor_a_dimen); @@ -693,7 +774,8 @@ void BatchNormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_buil model_builder, input, tensor_a_name, ANEURALNETWORKS_FUSED_NONE, - tensor_imm_product_name); + tensor_imm_product_name, + output_is_nhwc); // Add int32_t fuse_code = model_builder.FindActivation(output); @@ -701,7 +783,8 @@ void BatchNormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_buil model_builder, tensor_imm_product_name, tensor_b_name, fuse_code, - output); + output, + output_is_nhwc); } #pragma endregion op_batchnormalization @@ -782,17 +865,36 @@ bool PoolOpBuilder::IsOpSupportedImpl(ModelBuilder& model_builder, void PoolOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const ONNX_NAMESPACE::NodeProto& node) { auto& shaper(model_builder.GetShaper()); + const auto& operand_indices(model_builder.GetOperandIndices()); + const auto& operand_types(model_builder.GetOperandTypes()); + NodeAttrHelper helper(node); - const auto& input = node.input(0); + auto input = node.input(0); + bool use_nchw = model_builder.UseNCHW(); + bool input_is_nhwc = model_builder.IsOperandNHWC(input); + bool output_is_nhwc = false; + if (use_nchw) { + ORT_ENFORCE(!input_is_nhwc, "model_builder.UseNCHW() but input is NHWC"); + } else { + output_is_nhwc = true; + if (!input_is_nhwc) { + const auto& nchw_input = node.input(0); + if (!model_builder.GetNHWCOperand(nchw_input, input)) { + input = model_builder.GetUniqueName(nchw_input + "_nchw_to_nhwc"); + TransposeNCHWToNHWC(model_builder, nchw_input, input); + } + } + } + const auto& output = node.output(0); const auto& op = node.op_type(); - int32_t operationType; + int32_t op_type; if (op == "AveragePool" || op == "GlobalAveragePool") - operationType = ANEURALNETWORKS_AVERAGE_POOL_2D; + op_type = ANEURALNETWORKS_AVERAGE_POOL_2D; else // (op == "MaxPool" || op == "GlobalMaxPool") - operationType = ANEURALNETWORKS_MAX_POOL_2D; + op_type = ANEURALNETWORKS_MAX_POOL_2D; vector onnx_pads, onnx_strides, kernel_shape; if (op == "AveragePool" || op == "MaxPool") { @@ -811,12 +913,25 @@ void PoolOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, } int32_t fuse_code = model_builder.FindActivation(output); - AddPoolOperator(operationType, - model_builder, - input, - onnx_pads, onnx_strides, kernel_shape, - fuse_code, - output); + std::vector input_indices; + input_indices.push_back(operand_indices.at(input)); + input_indices.push_back(model_builder.AddOperandFromScalar(onnx_pads[1])); + input_indices.push_back(model_builder.AddOperandFromScalar(onnx_pads[3])); + input_indices.push_back(model_builder.AddOperandFromScalar(onnx_pads[0])); + input_indices.push_back(model_builder.AddOperandFromScalar(onnx_pads[2])); + input_indices.push_back(model_builder.AddOperandFromScalar(onnx_strides[1])); + input_indices.push_back(model_builder.AddOperandFromScalar(onnx_strides[0])); + input_indices.push_back(model_builder.AddOperandFromScalar(kernel_shape[1])); + input_indices.push_back(model_builder.AddOperandFromScalar(kernel_shape[0])); + input_indices.push_back(model_builder.AddOperandFromScalar(fuse_code)); + input_indices.push_back(model_builder.AddOperandFromScalar(use_nchw)); + + shaper.Pool(input, + onnx_pads, onnx_strides, kernel_shape, + use_nchw, + output); + const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); + model_builder.AddOperation(op_type, input_indices, {output}, {output_operand_type}, {output_is_nhwc}); } #pragma endregion op_pool @@ -877,7 +992,6 @@ void ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto& operand_types(model_builder.GetOperandTypes()); const auto& initializers(model_builder.GetInitializerTensors()); NodeAttrHelper helper(node); - bool use_nchw = model_builder.UseNCHW(); // onnx strides are in the order height, width // while nnapi strides are in the order width, height @@ -892,7 +1006,23 @@ void ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto onnx_dilations = helper.Get("dilations", vector{1, 1}); const auto group = helper.Get("group", 1); - const auto& input = node.input(0); + auto input = node.input(0); + bool use_nchw = model_builder.UseNCHW(); + bool input_is_nhwc = model_builder.IsOperandNHWC(input); + bool output_is_nhwc = false; + if (use_nchw) { + ORT_ENFORCE(!input_is_nhwc, "model_builder.UseNCHW() but input is NHWC"); + } else { + output_is_nhwc = true; + if (!input_is_nhwc) { + const auto& nchw_input = node.input(0); + if (!model_builder.GetNHWCOperand(nchw_input, input)) { + input = model_builder.GetUniqueName(nchw_input + "_nchw_to_nhwc"); + TransposeNCHWToNHWC(model_builder, nchw_input, input); + } + } + } + const auto& weight = node.input(1); const auto& output = node.output(0); @@ -905,7 +1035,7 @@ void ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, if (conv2d) { input_indices.push_back(AddInitializerInNewLayout( - model_builder, weight, L_NCHW)); + model_builder, weight, L_0231)); } else { // depthwise_conv2d input_indices.push_back(AddInitializerInNewLayout( model_builder, weight, L_1230)); @@ -927,13 +1057,13 @@ void ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto& weight_type = operand_types.at(weight).type; if (weight_type == Type::TENSOR_FLOAT32) { - float buffer[bias_dimen[0]]; - for (uint32_t i = 0; i < bias_dimen[0]; i++) { + vector buffer(bias_dimen[0]); + for (uint32_t i = 0; i < buffer.size(); i++) { buffer[i] = 0.f; } OperandType operandType(Type::TENSOR_FLOAT32, bias_dimen); bias_idx_val = model_builder.AddOperandFromPersistMemoryBuffer( - bias, &buffer[0], operandType); + bias, buffer.data(), operandType); } else { ORT_THROW("Unknown weight type " + TypeToStr(weight_type)); } @@ -952,7 +1082,7 @@ void ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, } int32_t fuse_code = model_builder.FindActivation(output); input_indices.push_back(model_builder.AddOperandFromScalar(fuse_code)); - // TODO support API 27 + // TODO support API 28 input_indices.push_back(model_builder.AddOperandFromScalar(use_nchw)); input_indices.push_back(model_builder.AddOperandFromScalar(onnx_dilations[1])); input_indices.push_back(model_builder.AddOperandFromScalar(onnx_dilations[0])); @@ -973,7 +1103,7 @@ void ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, } const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); - model_builder.AddOperation(operationCode, input_indices, {output}, {output_operand_type}); + model_builder.AddOperation(operationCode, input_indices, {output}, {output_operand_type}, {output_is_nhwc}); } #pragma endregion op_conv @@ -1017,6 +1147,8 @@ void CastOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto& input = node.input(0); const auto& output = node.output(0); + bool output_is_nhwc = model_builder.IsOperandNHWC(input); + auto to = helper.Get("to", 0); Type type; switch (to) { @@ -1035,7 +1167,8 @@ void CastOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, input_indices.push_back(operand_indices.at(input)); shaper.Identity(input, output); const OperandType output_operand_type(type, shaper[output]); - model_builder.AddOperation(ANEURALNETWORKS_CAST, input_indices, {output}, {output_operand_type}); + model_builder.AddOperation(ANEURALNETWORKS_CAST, input_indices, {output}, + {output_operand_type}, {output_is_nhwc}); } #pragma endregion @@ -1076,7 +1209,16 @@ void SoftMaxOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto& operand_types(model_builder.GetOperandTypes()); NodeAttrHelper helper(node); - const auto& input = node.input(0); + auto input = node.input(0); + if (model_builder.IsOperandNHWC(input)) { + // We want to transpose nhwc operand back to nchw before softmax + const auto& nhwc_input = node.input(0); + if (!model_builder.GetNCHWOperand(nhwc_input, input)) { + input = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw"); + TransposeNHWCToNCHW(model_builder, nhwc_input, input); + } + } + const auto& output = node.output(0); float beta = 1.f; int32_t axis = helper.Get("axis", 1); @@ -1087,7 +1229,8 @@ void SoftMaxOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, shaper.Identity(input, output); const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); - model_builder.AddOperation(ANEURALNETWORKS_SOFTMAX, input_indices, {output}, {output_operand_type}); + model_builder.AddOperation(ANEURALNETWORKS_SOFTMAX, input_indices, {output}, + {output_operand_type}, {false}); } #pragma endregion @@ -1110,12 +1253,14 @@ void IdentityOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto& input = node.input(0); const auto& output = node.output(0); + bool output_is_nhwc = model_builder.IsOperandNHWC(input); + std::vector input_indices; input_indices.push_back(operand_indices.at(input)); // input shaper.Identity(input, output); const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); - model_builder.RegisterOperand(output, operand_indices.at(input), output_operand_type); + model_builder.RegisterOperand(output, operand_indices.at(input), output_operand_type, output_is_nhwc); } #pragma endregion @@ -1141,6 +1286,16 @@ bool GemmOpBuilder::IsOpSupportedImpl( const auto& op = node.op_type(); const auto& initializers(model_builder.GetInitializerTensors()); + if (GetShape(model_builder.GetOnnxModel(), node.input(0)).size() != 2) { + LOGS_DEFAULT(VERBOSE) << "A must be 2D"; + return false; + } + + if (GetShape(model_builder.GetOnnxModel(), node.input(0)).size() != 2) { + LOGS_DEFAULT(VERBOSE) << "B must be 2D"; + return false; + } + if (op == "MatMul") { // Only support A*B B is an initializer if (!Contains(initializers, node.input(1))) { LOGS_DEFAULT(VERBOSE) << "B of MatMul must be known"; @@ -1172,7 +1327,10 @@ bool GemmOpBuilder::IsOpSupportedImpl( const auto c_shape = GetShape(model_builder.GetOnnxModel(), node.input(2)); if (c_shape.size() != 1 || c_shape[0] != (transB == 0 ? b_shape[1] : b_shape[0])) { - LOGS_DEFAULT(VERBOSE) << "C of Gemm must be a vector of b_shape[0]"; + LOGS_DEFAULT(VERBOSE) << "C of Gemm must be a vector of b_shape[0]" + << " b_shape: " << Shape2String(b_shape) + << " c_shape: " << Shape2String(c_shape); + return false; } } @@ -1243,8 +1401,8 @@ void GemmOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, shaper.FC(input1, input2, output); const OperandType output_operand_type(operand_types.at(input1).type, shaper[output]); - model_builder.AddOperation(ANEURALNETWORKS_FULLY_CONNECTED, - input_indices, {output}, {output_operand_type}); + model_builder.AddOperation(ANEURALNETWORKS_FULLY_CONNECTED, input_indices, {output}, + {output_operand_type}, {false}); } #pragma endregion @@ -1285,6 +1443,8 @@ void UnaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const auto& input = node.input(0); const auto& output = node.output(0); + bool output_is_nhwc = model_builder.IsOperandNHWC(input); + shaper.Identity(input, output); const OperandType output_operand_type(operand_types.at(input).type, shaper[output]); @@ -1312,7 +1472,97 @@ void UnaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, } std::vector input_indices; input_indices.push_back(operand_indices.at(input)); - model_builder.AddOperation(op_code, input_indices, {output}, {output_operand_type}); + model_builder.AddOperation(op_code, input_indices, {output}, {output_operand_type}, {output_is_nhwc}); +} + +#pragma endregion + +#pragma region op_concat + +class ConcatOpBuilder : public BaseOpBuilder { + private: + bool IsOpSupportedImpl(ModelBuilder& model_builder, + const ONNX_NAMESPACE::NodeProto& node) override; + + void AddToModelBuilderImpl(ModelBuilder& model_builder, + const ONNX_NAMESPACE::NodeProto& node) override; +}; + +bool ConcatOpBuilder::IsOpSupportedImpl(ModelBuilder& model_builder, + const ONNX_NAMESPACE::NodeProto& node) { + if (GetShape(model_builder.GetOnnxModel(), node.input(0)).size() > 4) { + LOGS_DEFAULT(VERBOSE) << "Concat supports at most 4D shape"; + return false; + } + + return true; +} + +void ConcatOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const ONNX_NAMESPACE::NodeProto& node) { + auto& shaper(model_builder.GetShaper()); + const auto& operand_indices(model_builder.GetOperandIndices()); + const auto& operand_types(model_builder.GetOperandTypes()); + NodeAttrHelper helper(node); + + std::vector input_indices; + const auto& input0 = node.input(0); + bool all_input_have_same_layout = true; + bool output_is_nhwc = false; + + // First we want to see if all the input are smae layout + for (int i = 0; i < node.input_size() - 1; i++) { + all_input_have_same_layout = + all_input_have_same_layout && + model_builder.IsOperandNHWC(node.input(i)) == model_builder.IsOperandNHWC(node.input(i + 1)); + } + + std::vector inputs; + inputs.reserve(node.input_size()); + if (all_input_have_same_layout) { + // if all the inputs are of same layout, output will be the same layout + if (model_builder.IsOperandNHWC(input0)) { + output_is_nhwc = true; + } + + for (const auto& input : node.input()) { + input_indices.push_back(operand_indices.at(input)); + inputs.push_back(input); + } + } else { + // if all the inputs are not same layout, + // will need transpos those nhwc tensors back to nchw + for (auto input : node.input()) { + if (model_builder.IsOperandNHWC(input)) { + std::string nhwc_input = input; + input = model_builder.GetUniqueName(input + "_nhwc_to_nchw"); + TransposeNHWCToNCHW(model_builder, nhwc_input, input); + } + input_indices.push_back(operand_indices.at(input)); + inputs.push_back(input); + } + } + + int32_t axis = helper.Get("axis", 1); + int rank = shaper[input0].size(); + if (axis < 0) { // NNAPI does not support negative axis + axis = rank + axis; + } + + if (output_is_nhwc) { + ORT_ENFORCE(rank == 4, "nhwc is only on 4d shape, input " + input0 + + " has rank: " + std::to_string(rank)); + // we are using nhwc here, but the axis is in nwhw, need to transpose axis from nchw to nhwc + const uint32_t axis_nchw_to_nhwc[4]{0, 3, 1, 2}; + axis = axis_nchw_to_nhwc[axis]; + } + input_indices.push_back(model_builder.AddOperandFromScalar(axis)); + + const auto& output = node.output(0); + shaper.Concat(inputs, axis, output); + const OperandType output_operand_type(operand_types.at(input0).type, shaper[output]); + model_builder.AddOperation(ANEURALNETWORKS_CONCATENATION, input_indices, {output}, + {output_operand_type}, {output_is_nhwc}); } #pragma endregion @@ -1368,6 +1618,8 @@ CreateOpBuilders() { op_map.emplace("Tanh", unary_op_builder); } + op_map.emplace("Concat", std::make_shared()); + return op_map; } diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/op_builder.h b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/op_builder.h index 128972768c..d0901ee1c1 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/op_builder.h +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/op_builder.h @@ -31,5 +31,9 @@ class IOpBuilder { std::unordered_map> CreateOpBuilders(); +void TransposeNHWCToNCHW(ModelBuilder& model_builder, + const std::string& input, + const std::string& output); + } // namespace nnapi -} // namespace onnxruntime \ No newline at end of file +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/shaper.cc b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/shaper.cc index ad2217aee6..3f9a2251d7 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/shaper.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/shaper.cc @@ -345,6 +345,39 @@ void Shaper::FC(const std::string& input1_name, const std::string& input2_name, } } +void Shaper::Concat(const std::vector& input_names, + const int32_t axis, + const std::string& output_name) { + std::vector dimens; + for (const auto& input_name : input_names) { + auto& dimen = shape_map_.at(input_name); + if (!dimens.empty()) { + for (size_t i = 0; i < dimens[0].size(); i++) { + if ((int32_t)i == axis) + continue; + + ORT_ENFORCE(dimen[i] == dimens[0][i], "Wrong input for concat"); + } + } + + dimens.push_back(shape_map_.at(input_name)); + } + + auto output_dimen = dimens[0]; + for (size_t i = 1; i < dimens.size(); i++) { + output_dimen[axis] += dimens[i][axis]; + } + + shape_map_[output_name] = output_dimen; + + if (!shaper_finalized_) { + shape_ops_.push_back( + [input_names, axis, output_name](Shaper& shaper) { + shaper.Concat(input_names, axis, output_name); + }); + } +} + void Shaper::AddShape(const std::string& name, const Shape& shape) { shape_map_[name] = shape; } @@ -354,10 +387,12 @@ void Shaper::UpdateShape(const std::string& name, const Shape& new_shape) { "Cannot UpdateShape while shaper is not finalized"); const auto& old_shape = shape_map_.at(name); - if (old_shape != new_shape && Product(shape_map_.at(name)) != 0) - ORT_THROW("The shape should be same size or old shape has size 0"); + if (old_shape != new_shape) { + if (Product(old_shape) != 0) + ORT_THROW("The shape should be same size or old shape has size 0 (dynamic shape)"); - shape_map_[name] = new_shape; + shape_map_[name] = new_shape; + } } void Shaper::UpdateDynamicDimensions() { @@ -373,3 +408,13 @@ void Shaper::Clear() { shape_map_.clear(); shape_ops_.clear(); } + +std::string Shape2String(const Shaper::Shape& shape) { + std::ostringstream os; + os << "[ "; + for (const auto& dim : shape) + os << dim << " "; + + os << "]"; + return os.str(); +} diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/shaper.h b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/shaper.h index 6a4818bcce..862c6134d9 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/shaper.h +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/shaper.h @@ -47,6 +47,10 @@ class Shaper { const std::string& input2_name, const std::string& output_name); + void Concat(const std::vector& input_names, + const int32_t axis, + const std::string& output_name); + // If the shape of certain input is dynamic // Use the following 2 functions to update the particular shape // and calculate the new output shape @@ -68,3 +72,5 @@ class Shaper { std::unordered_map shape_map_; std::vector> shape_ops_; }; + +std::string Shape2String(const Shaper::Shape& shape); diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/model.cc b/onnxruntime/core/providers/nnapi/nnapi_builtin/model.cc index 4fcd7b6bb3..a22dcf22cc 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/model.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/model.cc @@ -1,6 +1,8 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +#include + #include "model.h" #include "core/providers/nnapi/nnapi_builtin/builders/helper.h" #include "core/providers/nnapi/nnapi_builtin/nnapi_lib/nnapi_implementation.h" @@ -25,12 +27,18 @@ Model::~Model() { void Model::AddInput(const std::string& name, const android::nn::wrapper::OperandType& operand_type) { input_names_.push_back(name); - operand_types_.insert({name, operand_type}); + operand_types_.emplace(name, operand_type); } -void Model::AddOutput(const std::string& name, const android::nn::wrapper::OperandType& operand_type) { - output_names_.push_back(name); - operand_types_.insert({name, operand_type}); +void Model::AddOutput(const std::string& onnx_output_name, + const std::string& nnapi_output_name, + const android::nn::wrapper::OperandType& operand_type) { + LOGS_DEFAULT(VERBOSE) << "Model::AddOutput output name " << onnx_output_name + << " shape " << Shape2String(operand_type.dimensions); + + output_names_.push_back(onnx_output_name); + onnx_to_nnapi_output_map_.emplace(onnx_output_name, nnapi_output_name); + operand_types_.emplace(nnapi_output_name, operand_type); } const std::vector& Model::GetInputs() const { @@ -46,9 +54,10 @@ const android::nn::wrapper::OperandType& Model::GetInputType(const std::string& } const android::nn::wrapper::OperandType Model::GetOutputType(const std::string& name) const { - const auto& output_type = operand_types_.at(name); + const auto& nnapi_output_name = onnx_to_nnapi_output_map_.at(name); + const auto& output_type = operand_types_.at(nnapi_output_name); android::nn::wrapper::OperandType type( - output_type.type, shaper_for_exeuction_[name], output_type.operandType.scale, output_type.operandType.zeroPoint); + output_type.type, shaper_for_exeuction_[nnapi_output_name], output_type.operandType.scale, output_type.operandType.zeroPoint); return type; } @@ -79,6 +88,9 @@ void Model::SetInputBuffer(const int32_t index, const InputBuffer& input) { void Model::SetOutputBuffer(const int32_t index, const OutputBuffer& output) { PrepareForExecution(); + LOGS_DEFAULT(VERBOSE) << "Model::SetOutputBuffer, output shape " + << Shape2String(output.type.dimensions); + THROW_ON_ERROR(nnapi_->ANeuralNetworksExecution_setOutput( execution_, index, &output.type.operandType, output.buffer, output.type.GetOperandBlobByteSize())); } diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/model.h b/onnxruntime/core/providers/nnapi/nnapi_builtin/model.h index 985c02eb66..75c3fbb0f9 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/model.h +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/model.h @@ -4,6 +4,7 @@ #pragma once #include "builders/shaper.h" +#include "core/platform/ort_mutex.h" #include "nnapi_lib/NeuralNetworksWrapper.h" namespace onnxruntime { @@ -90,6 +91,9 @@ class Model { // Execute the NNAPI model void Predict(); + // Mutex for exclusive lock to this model object + OrtMutex& GetMutex() { return mutex_; } + private: const NnApi* nnapi_{nullptr}; bool prepared_for_exe_ = false; @@ -112,9 +116,21 @@ class Model { std::unordered_map input_map_; std::unordered_map output_map_; + // We may transpose the nnapi output to nchw with a different name + // This is map is to lookup the nnapi output from the onnx output + std::unordered_map onnx_to_nnapi_output_map_; + + OrtMutex mutex_; + Model(); void AddInput(const std::string& name, const android::nn::wrapper::OperandType& operand_type); - void AddOutput(const std::string& name, const android::nn::wrapper::OperandType& operand_type); + + // It is possible that the actual output from NNAPI model is not the same as the name of + // the output from the onnx model, need to have both names and add mapping between them + void AddOutput(const std::string& onnx_output_name, + const std::string& nnapi_output_name, + const android::nn::wrapper::OperandType& operand_type); + void SetShaper(const Shaper shaper) { shaper_ = shaper; } void SetInputBuffer(const int32_t index, const InputBuffer& input); diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.cc b/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.cc index 93b0c4f6ce..91201cb517 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.cc @@ -66,12 +66,12 @@ NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view onnxruntime::Graph& graph_build = model.MainGraph(); for (const auto& node : graph_view.Nodes()) { std::vector inputs, outputs; - for (auto input : node.InputDefs()) { + for (auto* input : node.InputDefs()) { auto& n_input = graph_build.GetOrCreateNodeArg(input->Name(), input->TypeAsProto()); inputs.push_back(&n_input); all_node_inputs.insert(input->Name()); } - for (auto output : node.OutputDefs()) { + for (auto* output : node.OutputDefs()) { auto& n_output = graph_build.GetOrCreateNodeArg(output->Name(), output->TypeAsProto()); outputs.push_back(&n_output); } @@ -113,9 +113,9 @@ NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view for (const auto& index : group) { sub_graph->nodes.push_back(node_index[index]); - const auto node = graph_view.GetNode(node_index[index]); + const auto* node = graph_view.GetNode(node_index[index]); - for (const auto& input : node->InputDefs()) { + for (const auto* input : node->InputDefs()) { const auto it = fused_outputs.find(input); if (it != fused_outputs.end()) { fused_outputs.erase(it); @@ -135,7 +135,7 @@ NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view std::unordered_set processed_outputs; for (auto it = node->OutputEdgesBegin(), end = node->OutputEdgesEnd(); it != end; ++it) { const auto node_idx = it->GetNode().Index(); - const auto output = node->OutputDefs()[it->GetSrcArgIndex()]; + const auto* output = node->OutputDefs()[it->GetSrcArgIndex()]; if (node_set.find(node_idx) != node_set.end()) { const auto iter = fused_inputs.find(output); @@ -152,7 +152,7 @@ NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view processed_outputs.insert(output); } - for (const auto& output : node->OutputDefs()) { + for (const auto* output : node->OutputDefs()) { if (processed_outputs.find(output) != processed_outputs.end()) continue; @@ -209,13 +209,6 @@ NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view return result; } -std::string GetShape(const std::vector& dimensions) { - std::string ret = ""; - for (auto dim : dimensions) - ret += std::to_string(dim) + " "; - return "[" + ret + "]"; -} - common::Status NnapiExecutionProvider::Compile(const std::vector& fused_nodes, std::vector& node_compute_funcs) { using namespace android::nn::wrapper; @@ -235,6 +228,8 @@ common::Status NnapiExecutionProvider::Compile(const std::vector nnapi_model = builder.Compile(); // Build map from input name to its index in input definitions @@ -275,8 +270,6 @@ common::Status NnapiExecutionProvider::Compile(const std::vector(state); const size_t num_inputs = ort.KernelContext_GetInputCount(context); const size_t num_outputs = ort.KernelContext_GetOutputCount(context); @@ -316,45 +309,50 @@ common::Status NnapiExecutionProvider::Compile(const std::vectorSetInputBuffers(inputs); - std::vector outputs; - outputs.reserve(num_outputs); - for (size_t i = 0; i < num_outputs; i++) { - const auto output_name = model->GetOutputs()[i]; - const auto model_output_type = model->GetOutputType(output_name); - const auto output_shape = model_output_type.dimensions; + // From this point we will need to take the exclusive lock on the model until the Predict is + // performed, to block other threads (if any) to modify this particular model + { + std::unique_lock lock(model->GetMutex()); + model->SetInputBuffers(inputs); + std::vector outputs; + outputs.reserve(num_outputs); + for (size_t i = 0; i < num_outputs; i++) { + const auto output_name = model->GetOutputs()[i]; + const auto model_output_type = model->GetOutputType(output_name); + const auto output_shape = model_output_type.dimensions; - std::vector int64_output_shape(output_shape.begin(), - output_shape.end()); - auto output_idx = model->GetMappedOutputIdx(output_name); - auto* output_tensor = ort.KernelContext_GetOutput(context, output_idx, - int64_output_shape.data(), - int64_output_shape.size()); + std::vector int64_output_shape(output_shape.begin(), + output_shape.end()); + auto output_idx = model->GetMappedOutputIdx(output_name); + auto* output_tensor = ort.KernelContext_GetOutput(context, output_idx, + int64_output_shape.data(), + int64_output_shape.size()); - void* output_buffer = nullptr; - switch (model_output_type.type) { - case Type::TENSOR_FLOAT32: - output_buffer = ort.GetTensorMutableData(output_tensor); - break; - case Type::TENSOR_INT32: - output_buffer = ort.GetTensorMutableData(output_tensor); - break; - default: - ORT_THROW("Unsupported output type: " + TypeToStr(model_output_type.type)); - break; + void* output_buffer = nullptr; + switch (model_output_type.type) { + case Type::TENSOR_FLOAT32: + output_buffer = ort.GetTensorMutableData(output_tensor); + break; + case Type::TENSOR_INT32: + output_buffer = ort.GetTensorMutableData(output_tensor); + break; + default: + return Status(common::ONNXRUNTIME, common::FAIL, + "Unsupported output type: " + TypeToStr(model_output_type.type)); + break; + } + + if (model_output_type.GetOperandBlobByteSize() == 0) { + return Status(common::ONNXRUNTIME, common::FAIL, "We do not support dynamic output shape for now"); + } + + outputs.push_back({output_buffer, std::move(model_output_type)}); } - if (model_output_type.GetOperandBlobByteSize() == 0) { - return Status(common::ONNXRUNTIME, common::FAIL, "We do not support dynamic output shape for now"); - } - - outputs.push_back({output_buffer, std::move(model_output_type)}); + model->SetOutputBuffers(outputs); + model->Predict(); } - model->SetOutputBuffers(outputs); - - model->Predict(); - return Status::OK(); }; diff --git a/onnxruntime/python/tools/transformers/benchmark_gpt2.py b/onnxruntime/python/tools/transformers/benchmark_gpt2.py index c1af86536f..a02fffb78c 100644 --- a/onnxruntime/python/tools/transformers/benchmark_gpt2.py +++ b/onnxruntime/python/tools/transformers/benchmark_gpt2.py @@ -75,7 +75,8 @@ def pytorch_inference(model, inputs, total_runs=100): input_ids, past, attention_mask, position_ids = inputs # Convert it back to fp32 as the PyTroch model cannot deal with half input. - attention_mask = attention_mask.to(dtype=torch.float32) if attention_mask else None + if attention_mask is not None: + attention_mask = attention_mask.to(dtype=torch.float32) past = [p.to(dtype=torch.float32) for p in past] latency = [] @@ -120,29 +121,37 @@ def onnxruntime_inference(ort_session, inputs, total_runs=100): return ort_outputs, average_latency -def get_dummy_inputs(batch_size, past_sequence_length, num_attention_heads, hidden_size, num_layer, vocab_size, device, - use_attention_mask, float16): +def get_dummy_inputs(batch_size, past_sequence_length, sequence_length, num_attention_heads, hidden_size, num_layer, + vocab_size, device, use_attention_mask, float16): float_type = torch.float16 if float16 else torch.float32 past_shape = [2, batch_size, num_attention_heads, past_sequence_length, int(hidden_size / num_attention_heads)] dummy_past = [torch.rand(past_shape, dtype=float_type, device=device) for _ in range(num_layer)] - dummy_input_ids = torch.randint(low=0, high=vocab_size - 1, size=(batch_size, 1), dtype=torch.int64, device=device) + dummy_input_ids = torch.randint(low=0, + high=vocab_size - 1, + size=(batch_size, sequence_length), + dtype=torch.int64, + device=device) if use_attention_mask: - dummy_attention_mask = torch.ones([batch_size, 1], dtype=float_type, device=device) - dummy_position_ids = torch.ones([batch_size, 1], dtype=torch.int64, device=device) * past_sequence_length + dummy_attention_mask = torch.ones([batch_size, past_sequence_length + sequence_length], + dtype=float_type, + device=device) + dummy_position_ids = torch.ones([batch_size, sequence_length], dtype=torch.int64, + device=device) * past_sequence_length return dummy_input_ids, dummy_past, dummy_attention_mask, dummy_position_ids return dummy_input_ids, dummy_past, None, None -def get_output_shapes(batch_size, past_sequence_length, config, use_LMHead): +def get_output_shapes(batch_size, past_sequence_length, sequence_length, config, use_LMHead): num_attention_heads = config.num_attention_heads hidden_size = config.hidden_size num_layer = config.n_layer vocab_size = config.vocab_size - last_state_shape = [batch_size, 1, vocab_size] if use_LMHead else [batch_size, 1, hidden_size] + last_state_shape = [batch_size, sequence_length, vocab_size + ] if use_LMHead else [batch_size, sequence_length, hidden_size] present_state_shape = [ - 2, batch_size, num_attention_heads, past_sequence_length + 1, + 2, batch_size, num_attention_heads, past_sequence_length + sequence_length, int(hidden_size / num_attention_heads) ] @@ -327,7 +336,7 @@ def parse_arguments(): parser.add_argument('-b', '--batch_sizes', nargs='+', type=int, default=[1], help="batch size") parser.add_argument('-s', - '--sequence_lengths', + '--past_sequence_lengths', nargs='+', type=int, default=[8, 16, 32, 64, 128, 256], @@ -375,23 +384,24 @@ def export_onnx(model, config, tokenizer, device, output_dir, use_LMHead, use_at # Shape of input tensors: # input_ids: (batch_size, seq_len) # past_{i}: (2, batch_size, num_heads, past_seq_len, hidden_size/num_heads) - # attention_mask: (batch_size, seq_len) + # attention_mask: (batch_size, past_seq_len + seq_len) # Shape of output tensors: # last_state: (batch_size, seq_len, hidden_size) # or prediction_scores: (batch_size, seq_len, vocab_size) - # present_{i}: (2, batch_size, num_heads, all_seq_len, hidden_size/num_heads) + # present_{i}: (2, batch_size, num_heads, past_seq_len + seq_len, hidden_size/num_heads) dynamic_axes = {'input_ids': {0: 'batch_size', 1: 'seq_len'}, output_names[0]: {0: 'batch_size', 1: 'seq_len'}} for name in past_names: dynamic_axes[name] = {1: 'batch_size', 3: 'past_seq_len'} for name in present_names: - dynamic_axes[name] = {1: 'batch_size', 3: 'all_seq_len'} + dynamic_axes[name] = {1: 'batch_size', 3: 'past_seq_len+seq_len'} if use_attention_mask: - dynamic_axes['attention_mask'] = {0: 'batch_size', 1: 'all_seq_len'} + dynamic_axes['attention_mask'] = {0: 'batch_size', 1: 'past_seq_len+seq_len'} dynamic_axes['position_ids'] = {0: 'batch_size', 1: 'seq_len'} dummy_inputs = get_dummy_inputs(batch_size=1, past_sequence_length=1, + sequence_length=1, num_attention_heads=config.num_attention_heads, hidden_size=config.hidden_size, num_layer=num_layer, @@ -469,17 +479,15 @@ def main(): if not os.path.exists(output_dir): os.makedirs(output_dir) - use_torchscript = False - model_class = MyGPT2LMHeadModel if args.model_class == 'GPT2LMHeadModel' else MyGPT2Model use_LMHead = (args.model_class == 'GPT2LMHeadModel') model_name = args.model_name - config = AutoConfig.from_pretrained(model_name, torchscript=use_torchscript, cache_dir=cache_dir) + config = AutoConfig.from_pretrained(model_name, torchscript=False, cache_dir=cache_dir) model = model_class.from_pretrained(model_name, config=config, cache_dir=cache_dir) tokenizer = GPT2Tokenizer.from_pretrained(model_name, cache_dir=cache_dir) - #if use_torchscript: - # model = torch.jit.trace(model, (input_ids, past)) + + # This scirpt does not support float16 for PyTorch. #if args.float16: # model.half() @@ -521,10 +529,14 @@ def main(): if session is None: return + # One word is generated for each inference. This length does not include that of past state. + sequence_length = 1 + # Allocate output buffers for IO Binding output_buffers = {} if not args.disable_ort_io_binding: - max_output_shapes = get_output_shapes(max(args.batch_sizes), max(args.sequence_lengths), config, use_LMHead) + max_output_shapes = get_output_shapes(max(args.batch_sizes), max(args.past_sequence_lengths), sequence_length, + config, use_LMHead) output_buffers = get_output_buffers(max_output_shapes, device, args.float16) csv_filename = args.result_csv or "benchmark_result_{}.csv".format(datetime.now().strftime("%Y%m%d-%H%M%S")) @@ -537,12 +549,12 @@ def main(): csv_writer.writeheader() for batch_size in args.batch_sizes: - for past_sequence_length in args.sequence_lengths: + for past_sequence_length in args.past_sequence_lengths: logger.debug(f"Running test for batch_size={batch_size} past_sequence_length={past_sequence_length}...") - dummy_inputs = get_dummy_inputs(batch_size, past_sequence_length, config.num_attention_heads, - config.hidden_size, config.n_layer, config.vocab_size, device, - args.use_attention_mask, args.float16) - output_shapes = get_output_shapes(batch_size, past_sequence_length, config, use_LMHead) + dummy_inputs = get_dummy_inputs(batch_size, past_sequence_length, sequence_length, + config.num_attention_heads, config.hidden_size, config.n_layer, + config.vocab_size, device, args.use_attention_mask, args.float16) + output_shapes = get_output_shapes(batch_size, past_sequence_length, sequence_length, config, use_LMHead) try: latencies = inference(model, diff --git a/onnxruntime/python/tools/transformers/fusion_gpt_attention.py b/onnxruntime/python/tools/transformers/fusion_gpt_attention.py index 31d7b1ebff..9089f8c5ad 100644 --- a/onnxruntime/python/tools/transformers/fusion_gpt_attention.py +++ b/onnxruntime/python/tools/transformers/fusion_gpt_attention.py @@ -20,12 +20,13 @@ class FusionGptAttention(Fusion): def __init__(self, model: OnnxModel, num_heads: int): super().__init__(model, "Attention", "LayerNormalization", "with past") self.num_heads = num_heads + self.utils = FusionUtils(model) + self.casted_attention_mask = {} # map from name of attention mask to the name that casted to int32 - def create_attention_node(self, gemm, gemm_qkv, past, present, input, output): + def create_attention_node(self, gemm, gemm_qkv, past, present, input, output, mask=''): attention_node_name = self.model.create_node_name('GptAttention') - mask_index = '' attention_node = helper.make_node('Attention', - inputs=[input, gemm.input[1], gemm.input[2], mask_index, past], + inputs=[input, gemm.input[1], gemm.input[2], mask, past], outputs=[attention_node_name + "_output", present], name=attention_node_name) attention_node.domain = "com.microsoft" @@ -114,6 +115,7 @@ class FusionGptAttention(Fusion): logger.debug("Add and LayerNormalization shall have one same input") return + input_mask_nodes = None qk_nodes = self.model.match_parent_path(matmul_qkv, ['Softmax', 'Sub', 'Mul', 'Div', 'MatMul'], [0, 0, 0, 0, 0]) if qk_nodes is not None: (softmax_qk, sub_qk, mul_qk, div_qk, matmul_qk) = qk_nodes @@ -122,7 +124,7 @@ class FusionGptAttention(Fusion): ['Mul', 'Sub', 'Slice', 'Slice', 'Unsqueeze', 'Sub', 'Squeeze', 'Slice', 'Shape', 'Div'], [1, 0, 1, 0, 1, 0, 0, 0, 0, 0]) # yapf: disable if mask_nodes is None: - logger.debug("fuse_attention: failed to match mask path") + logger.debug("fuse_attention: failed to match unidirectional mask path") return div_mask = mask_nodes[-1] @@ -131,11 +133,27 @@ class FusionGptAttention(Fusion): return else: # New pattern for gpt2 from PyTorch 1.5.0 and Transformers 2.9.0. - qk_nodes = self.model.match_parent_path(matmul_qkv, ['Softmax', 'Where', 'Div', 'MatMul'], [0, 0, 1, 0]) + i, qk_nodes, _ = self.model.match_parent_paths( + matmul_qkv, [(['Softmax', 'Where', 'Div', 'MatMul'], [0, 0, 1, 0]), + (['Softmax', 'Add', 'Where', 'Div', 'MatMul'], [0, 0, 0, 1, 0])], output_name_to_node) if qk_nodes is None: - logger.debug("fuse_attention: failed to match qk path") + logger.debug("fuse_attention: failed to match qk nodes") return - (softmax_qk, where_qk, div_qk, matmul_qk) = qk_nodes + + where_qk = qk_nodes[-3] + div_qk = qk_nodes[-2] + matmul_qk = qk_nodes[-1] + + if i == 1: + add_qk = qk_nodes[1] + _, input_mask_nodes, _ = self.model.match_parent_paths( + add_qk, [(['Mul', 'Sub', 'Cast', 'Unsqueeze', 'Unsqueeze', 'Reshape'], [1, 0, 1, 0, 0, 0]), + (['Mul', 'Sub', 'Unsqueeze', 'Unsqueeze', 'Reshape'], [1, 0, 1, 0, 0])], + output_name_to_node) + if input_mask_nodes is None: + logger.debug("fuse_attention: failed to match input attention mask path") + return + mask_nodes = self.model.match_parent_path( where_qk, ['Cast', 'Slice', 'Slice', 'Unsqueeze', 'Sub', 'Squeeze', 'Slice', 'Shape', 'Div'], @@ -188,8 +206,20 @@ class FusionGptAttention(Fusion): logger.info("expect past to be same") return + attention_mask_input_name = '' + if input_mask_nodes is not None: + input_name = input_mask_nodes[-1].input[0] + if input_name in self.casted_attention_mask: + attention_mask_input_name = self.casted_attention_mask[input_name] + elif self.model.find_graph_input(input_name): + casted, attention_mask_input_name = self.utils.cast_graph_input_to_int32(input_name) + self.casted_attention_mask[input_name] = attention_mask_input_name + else: + attention_mask_input_name, cast_node = self.utils.cast_input_to_int32(input_name) + self.casted_attention_mask[input_name] = attention_mask_input_name + self.create_attention_node(gemm, gemm_qkv, past, present, layernorm_before_attention.output[0], - reshape_qkv.output[0]) + reshape_qkv.output[0], attention_mask_input_name) # we rely on prune_graph() to clean old subgraph nodes: # qk_nodes + q_nodes + k_nodes + v_nodes + mask_nodes + [reshape_qkv, transpose_qkv, matmul_qkv] diff --git a/onnxruntime/test/contrib_ops/attention_op_test.cc b/onnxruntime/test/contrib_ops/attention_op_test.cc index 897401c197..831505ad0e 100644 --- a/onnxruntime/test/contrib_ops/attention_op_test.cc +++ b/onnxruntime/test/contrib_ops/attention_op_test.cc @@ -1,4 +1,5 @@ // Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. #include "gtest/gtest.h" @@ -8,12 +9,17 @@ namespace onnxruntime { namespace test { +enum MaskIndexType { + kMaskIndexEnd = 0, + kMaskIndexEndAndStart, + kMaskRaw +}; static void RunAttentionTest( const std::vector& input_data, // input: [batch_size, sequence_length, hidden_size] const std::vector& weights_data, // weights: [hidden_size, 3 * hidden_size] const std::vector& bias_data, // bias: [3 * hidden_size] - const std::vector& mask_index_data, // mask_index: [batch_size] or empty + const std::vector& mask_index_data, // mask_index: [batch_size] or [batch_size, past_sequence_length + sequence_length] or empty const std::vector& output_data, // output: [batch_size, sequence_length, hidden_size] int batch_size, int sequence_length, @@ -23,14 +29,14 @@ static void RunAttentionTest( bool is_unidirectional = false, bool use_past_state = false, int past_sequence_length = 0, - int head_size = 0, const std::vector* past_data = nullptr, - const std::vector* present_data = nullptr) { + const std::vector* present_data = nullptr, + MaskIndexType mask_index_type = kMaskIndexEnd) { int min_cuda_architecture = use_float16 ? 530 : 0; bool enable_cuda = HasCudaEnvironment(min_cuda_architecture); bool enable_cpu = (nullptr != DefaultCpuExecutionProvider().get()) && !use_float16; - + int head_size = hidden_size / number_of_heads; if (enable_cpu || enable_cuda) { OpTester tester("Attention", 1, onnxruntime::kMSDomain); tester.AddAttribute("num_heads", static_cast(number_of_heads)); @@ -39,9 +45,14 @@ static void RunAttentionTest( std::vector input_dims = {batch_size, sequence_length, hidden_size}; std::vector weights_dims = {hidden_size, 3 * hidden_size}; std::vector bias_dims = {3 * hidden_size}; - std::vector mask_index_dims = {batch_size}; - std::vector past_dims = {2, batch_size, head_size, past_sequence_length, head_size}; - std::vector present_dims = {2, batch_size, head_size, past_sequence_length + sequence_length, head_size}; + + std::vector mask_index_dims_1 = {batch_size}; + std::vector mask_index_dims_2 = {2 * batch_size}; + std::vector mask_index_dims_3 = {batch_size, past_sequence_length + sequence_length}; + std::vector mask_index_dims = (mask_index_type == kMaskIndexEnd ? mask_index_dims_1 : (mask_index_type == kMaskIndexEndAndStart ? mask_index_dims_2 : mask_index_dims_3)); + + std::vector past_dims = {2, batch_size, number_of_heads, past_sequence_length, head_size}; + std::vector present_dims = {2, batch_size, number_of_heads, past_sequence_length + sequence_length, head_size}; std::vector output_dims = input_dims; if (use_float16) { @@ -59,8 +70,7 @@ static void RunAttentionTest( if (mask_index_data.size() > 0) { // mask index is optional. tester.AddInput("mask_index", mask_index_dims, mask_index_data); } else { - std::vector dims = {static_cast(mask_index_data.size())}; - tester.AddInput("", dims, mask_index_data); + tester.AddMissingOptionalInput(); } if (use_past_state) { @@ -451,10 +461,10 @@ TEST(AttentionTest, AttentionEmptyPastState) { bool is_unidirectional = true; bool use_past_state = true; int past_sequence_length = 0; - int head_size = 2; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional, - use_past_state, past_sequence_length, head_size, &past_data, &present_data); + use_past_state, past_sequence_length, &past_data, &present_data); } TEST(AttentionTest, AttentionPastStateBatch1) { @@ -550,10 +560,10 @@ TEST(AttentionTest, AttentionPastStateBatch1) { bool is_unidirectional = true; bool use_past_state = true; int past_sequence_length = 3; - int head_size = 2; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional, - use_past_state, past_sequence_length, head_size, &past_data, &present_data); + use_past_state, past_sequence_length, &past_data, &present_data); } TEST(AttentionTest, AttentionPastStateBatch2) { @@ -653,10 +663,516 @@ TEST(AttentionTest, AttentionPastStateBatch2) { bool is_unidirectional = true; bool use_past_state = true; int past_sequence_length = 3; - int head_size = 2; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional, - use_past_state, past_sequence_length, head_size, &past_data, &present_data); + use_past_state, past_sequence_length, &past_data, &present_data); +} + +TEST(AttentionTest, AttentionPastStateBatch2WithPadding) { + int batch_size = 2; + int sequence_length = 1; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + -0.10902753f, 0.0041178204f, 0.1871525f, -0.20399982f, + 0.027207348f, -0.25321805f, 0.12869114f, 0.023136809f}; + + std::vector weight_data = { + -0.4738484025001526f, + -0.2613658607006073f, + -0.0978037416934967f, + -0.34988933801651f, + 0.2243240624666214f, + -0.0429205559194088f, + 0.418695330619812f, + 0.17441125214099884f, + -0.18825532495975494f, + 0.18357256054878235f, + -0.5806483626365662f, + -0.02251487597823143f, + + 0.08742205798625946f, + 0.14734269678592682f, + 0.2387014478445053f, + 0.2884027063846588f, + 0.6490834355354309f, + 0.16965825855731964f, + -0.06346885114908218f, + 0.4073973298072815f, + -0.03070945478975773f, + 0.4110257923603058f, + 0.07896808534860611f, + 0.16783113777637482f, + + 0.0038893644232302904f, + 0.06946629285812378f, + 0.36680519580841064f, + -0.07261059433221817f, + -0.14960581064224243f, + 0.020944256335496902f, + -0.09378612786531448f, + -0.1336742341518402f, + 0.06061394885182381f, + 0.2205914407968521f, + -0.03519909828901291f, + -0.18405692279338837f, + + 0.22149960696697235f, + -0.1884360909461975f, + -0.014074507169425488f, + 0.4252440333366394f, + 0.24987126886844635f, + -0.31396418809890747f, + 0.14036843180656433f, + 0.2854192554950714f, + 0.09709841012954712f, + 0.09935075044631958f, + -0.012154420837759972f, + 0.2575816512107849f}; + + std::vector bias_data = { + 0.4803391396999359f, + -0.5254325866699219f, + -0.42926454544067383f, + -0.2059524953365326f, + -0.12773379683494568f, + -0.09542735666036606f, + -0.35286077857017517f, + -0.07646317780017853f, + -0.04590314254164696f, + -0.03752850368618965f, + -0.013764488510787487f, + -0.18478283286094666f}; + + // One sequence has both left padding and right padding + std::vector mask_index_data = {4, 3, 0, 2}; + + std::vector output_data = { + 0.14902574f, 0.62273371f, 0.43022552f, 0.12759127f, + 0.18029204f, 0.07451740f, 0.73694098f, 0.17766341f}; + + std::vector past_data = { + 0.42028648f, 0.55855948f, 0.044569403f, 0.76525789f, 0.13962431f, 0.40977913f, 0.36911047f, 0.83399564f, 0.36905321f, 0.91414654f, 0.17300875f, 0.78793788f, + 0.10279467f, 0.80501258f, 0.089550517f, 0.85371113f, 0.61801594f, 0.91222942f, 0.88626182f, 0.069776468f, 0.10591964f, 0.84836882f, 0.83520192f, 0.0098680854f, + 0.3113814f, 0.63999802f, 0.28603253f, 0.98899829f, 0.044405211f, 0.95105386f, 0.81278932f, 0.63969064f, 0.14494057f, 0.11349615f, 0.87086016f, 0.20983537f, + 0.35107401f, 0.90144604f, 0.68950737f, 0.18928574f, 0.18029204f, 0.074517399f, 0.70763874f, 0.48440042f, 0.58114725f, 0.1048766f, 0.73694098f, 0.17766342f}; + + std::vector present_data = { + 0.42028648f, 0.55855948f, 0.044569403f, 0.76525789f, 0.13962431f, 0.40977913f, -0.22849128f, -0.022080801f, 0.36911047f, 0.83399564f, 0.36905321f, 0.91414654f, 0.17300875f, 0.78793788f, -0.4449589f, -0.17704415f, 0.10279467f, 0.80501258f, 0.089550517f, 0.85371113f, 0.61801594f, 0.91222942f, -0.2994619f, -0.14412443f, 0.88626182f, 0.069776468f, 0.10591964f, 0.84836882f, 0.83520192f, 0.0098680854f, -0.33421949f, -0.18547727f, + 0.3113814f, 0.63999802f, 0.28603253f, 0.98899829f, 0.044405211f, 0.95105386f, -0.033968594f, -0.034833729f, 0.81278932f, 0.63969064f, 0.14494057f, 0.11349615f, 0.87086016f, 0.20983537f, 0.045759238f, -0.26863033f, 0.35107401f, 0.90144604f, 0.68950737f, 0.18928574f, 0.18029204f, 0.074517399f, -0.033201858f, -0.10592631f, 0.70763874f, 0.48440042f, 0.58114725f, 0.1048766f, 0.73694098f, 0.17766342f, -0.054369561f, -0.24562015f}; + + bool is_unidirectional = true; + bool use_past_state = true; + int past_sequence_length = 3; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional, + use_past_state, past_sequence_length, &past_data, &present_data, kMaskIndexEndAndStart); +} + +TEST(AttentionTest, AttentionBatch2MaskIndex2) { + int batch_size = 2; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + std::vector mask_index_data = {2, 2, 0, 0}; + + std::vector output_data = { + 3.1495983600616455f, 0.10843668878078461f, 4.25f, 5.6499996185302734f, + 3.9696791172027588f, 0.073143675923347473f, 4.2499995231628418f, 5.6499991416931152f, + 3.1495983600616455f, 0.10843668878078461f, 4.25f, 5.6499996185302734f, + 3.9696791172027588f, 0.073143675923347473f, 4.2499995231628418f, 5.6499991416931152f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskIndexEndAndStart); +} + +TEST(AttentionTest, AttentionRightPaddingMaskIndex2) { + int batch_size = 1; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test mask_index < sequence_length + std::vector mask_index_data = {1, 0}; + + std::vector output_data = { + 8.6899995803833008f, -0.13000002503395081f, 4.25f, 5.6499996185302734f, + 8.6899995803833008f, -0.13000002503395081f, 4.2499995231628418f, 5.6499991416931152f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskIndexEndAndStart); +} + +TEST(AttentionTest, AttentionLeftPaddingMaskIndex2) { + int batch_size = 1; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test mask start position > 0. + std::vector mask_index_data = {2, 1}; + + std::vector output_data = { + 8.69f, -0.13f, 4.25f, 5.65f, + 8.69f, -0.13f, 4.25f, 5.65f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskIndexEndAndStart); +} + +TEST(AttentionTest, AttentionBatch2LeftPaddingMaskIndex2) { + int batch_size = 2; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test mask start position > 0. + std::vector mask_index_data = {2, 2, 1, 0}; + + std::vector output_data = { + 8.69f, -0.13f, 4.25f, 5.65f, + 8.69f, -0.13f, 4.25f, 5.65f, + 3.14959716796875f, 0.10843672603368759f, 4.25f, 5.65f, + 3.9696791172027588f, 0.073143675923347473f, 4.25f, 5.65f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskIndexEndAndStart); +} + +TEST(AttentionTest, AttentionBatch2AttentionMask) { + int batch_size = 2; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test mask start position > 0. + std::vector mask_index_data = {0, 1, 1, 1}; + + std::vector output_data = { + 8.69f, -0.13f, 4.25f, 5.65f, + 8.69f, -0.13f, 4.25f, 5.65f, + 3.14959716796875f, 0.10843672603368759f, 4.25f, 5.65f, + 3.9696791172027588f, 0.073143675923347473f, 4.25f, 5.65f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskRaw); +} + +TEST(AttentionTest, AttentionUnidirectionalAttentionMask) { + int batch_size = 2; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test mask start position > 0. + std::vector mask_index_data = {0, 1, 1, 1}; + + std::vector output_data = { + 3.967245340f, 0.07324841f, 4.25f, 5.65f, + 8.69f, -0.13f, 4.25f, 5.65f, + 8.69f, -0.13f, 4.25f, 5.65f, + 3.96967912f, 0.07314367f, 4.25f, 5.65f}; + + bool use_float16 = false; + bool is_unidirectional = true; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskRaw); +} + +TEST(AttentionTest, AttentionMask1DEndNoWord) { + int batch_size = 2; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test that all attention masks are zero. + std::vector mask_index_data = {0, 0}; + + std::vector output_data = { + 3.96724534f, 0.07324841f, 4.25f, 5.65f, + 3.14984703f, 0.10842596f, 4.25f, 5.65f, + 3.14984703f, 0.10842596f, 4.25f, 5.65f, + 3.96724534f, 0.07324841f, 4.25f, 5.65f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskIndexEnd); +} + +TEST(AttentionTest, AttentionMask1DNoWord) { + int batch_size = 2; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test that all attention masks are zero. + std::vector mask_index_data = {0, 0, 2, 2}; + + std::vector output_data = { + 3.96724534f, 0.07324841f, 4.25f, 5.65f, + 3.14984703f, 0.10842596f, 4.25f, 5.65f, + 3.14984703f, 0.10842596f, 4.25f, 5.65f, + 3.96724534f, 0.07324841f, 4.25f, 5.65f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskIndexEndAndStart); +} + +TEST(AttentionTest, AttentionMask2DNoWord) { + int batch_size = 2; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test that all attention masks are zero. + std::vector mask_index_data = {0, 0, 0, 0}; + + std::vector output_data = { + 3.96724534f, 0.07324841f, 4.25f, 5.65f, + 3.14984703f, 0.10842596f, 4.25f, 5.65f, + 3.14984703f, 0.10842596f, 4.25f, 5.65f, + 3.96724534f, 0.07324841f, 4.25f, 5.65f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskRaw); +} + +TEST(AttentionTest, AttentionMaskIndexOutOfRange) { + int batch_size = 2; + int sequence_length = 2; + int hidden_size = 4; + int number_of_heads = 2; + + std::vector input_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f, + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f, + 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f, + 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f, + 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f}; + + // Test end_position > sequence length, or start_position < 0 + std::vector mask_index_data = {3, 2, 0, -1}; + + std::vector output_data = { + 3.1495983600616455f, 0.10843668878078461f, 4.25f, 5.6499996185302734f, + 3.9696791172027588f, 0.073143675923347473f, 4.2499995231628418f, 5.6499991416931152f, + 3.1495983600616455f, 0.10843668878078461f, 4.25f, 5.6499996185302734f, + 3.9696791172027588f, 0.073143675923347473f, 4.2499995231628418f, 5.6499991416931152f}; + + bool use_float16 = false; + bool is_unidirectional = false; + bool use_past_state = false; + int past_sequence_length = 0; + const std::vector* past_data = nullptr; + const std::vector* present_data = nullptr; + RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data, + batch_size, sequence_length, hidden_size, number_of_heads, + use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data, kMaskIndexEndAndStart); } TEST(AttentionTest, AttentionPastState_dynamic) { diff --git a/onnxruntime/test/onnx/main.cc b/onnxruntime/test/onnx/main.cc index 3ae3ebee4b..db82906abe 100644 --- a/onnxruntime/test/onnx/main.cc +++ b/onnxruntime/test/onnx/main.cc @@ -433,7 +433,7 @@ int real_main(int argc, char* argv[], Ort::Env& env) { static const ORTCHAR_T* cuda_flaky_tests[] = { ORT_TSTR("fp16_inception_v1"), ORT_TSTR("fp16_shufflenet"), ORT_TSTR("fp16_tiny_yolov2")}; - static const ORTCHAR_T* dml_disabled_tests[] = {ORT_TSTR("mlperf_ssd_resnet34_1200"), ORT_TSTR("mlperf_ssd_mobilenet_300"), ORT_TSTR("mask_rcnn"), ORT_TSTR("faster_rcnn"), ORT_TSTR("tf_pnasnet_large"), ORT_TSTR("zfnet512")}; + static const ORTCHAR_T* dml_disabled_tests[] = {ORT_TSTR("mlperf_ssd_resnet34_1200"), ORT_TSTR("mlperf_ssd_mobilenet_300"), ORT_TSTR("mask_rcnn"), ORT_TSTR("faster_rcnn"), ORT_TSTR("tf_pnasnet_large"), ORT_TSTR("zfnet512"), ORT_TSTR("keras2coreml_Dense_ImageNet")}; static const ORTCHAR_T* dnnl_disabled_tests[] = {ORT_TSTR("test_densenet121"), ORT_TSTR("test_resnet18v2"), ORT_TSTR("test_resnet34v2"), ORT_TSTR("test_resnet50v2"), ORT_TSTR("test_resnet101v2"), ORT_TSTR("test_resnet101v2"), ORT_TSTR("test_vgg19"), ORT_TSTR("tf_inception_resnet_v2"), ORT_TSTR("tf_inception_v1"), ORT_TSTR("tf_inception_v3"), ORT_TSTR("tf_inception_v4"), ORT_TSTR("tf_mobilenet_v1_1.0_224"), ORT_TSTR("tf_mobilenet_v2_1.0_224"), ORT_TSTR("tf_mobilenet_v2_1.4_224"), ORT_TSTR("tf_nasnet_large"), ORT_TSTR("tf_pnasnet_large"), ORT_TSTR("tf_resnet_v1_50"), ORT_TSTR("tf_resnet_v1_101"), ORT_TSTR("tf_resnet_v1_101"), diff --git a/onnxruntime/test/optimizer/nchwc_optimizer_test.cc b/onnxruntime/test/optimizer/nchwc_optimizer_test.cc index f51cc59205..64ba6f5e2a 100644 --- a/onnxruntime/test/optimizer/nchwc_optimizer_test.cc +++ b/onnxruntime/test/optimizer/nchwc_optimizer_test.cc @@ -1219,6 +1219,41 @@ TEST(NchwcOptimizerTests, Upsample) { } } +TEST(NchwcOptimizerTests, Activation) { + auto test_case = [&](const std::string& activation_op_type) { + auto build_test_case = [&](NchwcTestHelper& helper) { + auto* input_arg = helper.MakeInput({1, 48, 11, 15}); + auto* conv1_output_arg = helper.MakeIntermediate(); + auto* activation_output_arg = helper.MakeIntermediate(); + auto* mul_output_arg = helper.MakeIntermediate(); + auto* output_arg = helper.MakeOutput(); + + helper.AddConvNode(input_arg, conv1_output_arg, {32, 48, 3, 3}); + helper.AddNode(activation_op_type, {conv1_output_arg}, {activation_output_arg}); + helper.AddNode("Add", {conv1_output_arg, activation_output_arg}, {mul_output_arg}); + helper.AddConvNode(mul_output_arg, output_arg, {16, 32, 1, 1}); + }; + + auto check_nchwc_graph = [&](NchwcInferenceSession& session) { + auto op_to_count = session.CountOpsInGraph(); + EXPECT_EQ(op_to_count["nchwc.Conv"], 2); + EXPECT_EQ(op_to_count["nchwc.ReorderInput"], 1); + EXPECT_EQ(op_to_count["nchwc.ReorderOutput"], 1); + EXPECT_EQ(op_to_count[activation_op_type], 1); + EXPECT_EQ(op_to_count["Add"], 1); + }; + + NchwcOptimizerTester(build_test_case, check_nchwc_graph); + }; + + // Verify that the optimizer doesn't add reorders for these activations that + // cannot be fused with a convolution. + std::vector activation_op_types{"Relu", "Sigmoid", "Tanh"}; + for (auto& activation_op_type : activation_op_types) { + test_case(activation_op_type); + } +} + #endif } // namespace test diff --git a/onnxruntime/test/providers/cpu/math/matmul_test.cc b/onnxruntime/test/providers/cpu/math/matmul_test.cc index 1e6bece98c..418dee24f0 100644 --- a/onnxruntime/test/providers/cpu/math/matmul_test.cc +++ b/onnxruntime/test/providers/cpu/math/matmul_test.cc @@ -77,6 +77,13 @@ std::vector> GenerateTestCases() {2, 2, 4}, {20, 23, 26, 29, 56, 68, 80, 92, 92, 113, 134, 155, 128, 158, 188, 218}}); + test_cases.push_back( + {"test 2D special 3", + {2, 6}, + {1, 1, 6, 1}, + {1, 1, 2, 1}, + {55, 145}}); + return test_cases; } diff --git a/orttraining/tools/ci_test/run_bert_perf_test.py b/orttraining/tools/ci_test/run_bert_perf_test.py index 14f72a2c92..d3909f36f5 100644 --- a/orttraining/tools/ci_test/run_bert_perf_test.py +++ b/orttraining/tools/ci_test/run_bert_perf_test.py @@ -26,7 +26,7 @@ def main(): Config = namedtuple('Config', ['use_mixed_precision', 'max_seq_length', 'batch_size', 'max_predictions_per_seq']) configs = [ - Config(True, 128, 66, 20), + Config(True, 128, 64, 20), Config(True, 512, 10, 80), Config(False, 128, 33, 20), Config(False, 512, 5, 80)