mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Update gather to use multiple threads (#11524)
This commit is contained in:
parent
5eaa893936
commit
deef214772
2 changed files with 71 additions and 88 deletions
|
|
@ -1,6 +1,7 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include <string>
|
||||
#include "gather_elements.h"
|
||||
#include "onnxruntime_config.h"
|
||||
|
||||
|
|
@ -25,65 +26,44 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
DataTypeImpl::GetTensorType<int64_t>()}),
|
||||
GatherElements);
|
||||
|
||||
// Some helpers needed for GatherElements op -
|
||||
|
||||
// The following method computes the offset in the flattened array
|
||||
// using every axis except the inner dimension (as the offset is just 1)
|
||||
// and the axis that 'GatherElements' is processing for as that requires the corresponding
|
||||
// 'indices' value
|
||||
// This prevents the need to compute this offset for every element within the same 'inner_dimension' chunk
|
||||
// as this value just differs by 1 for the chunk elements and we can have this cached and re-use as needed
|
||||
static inline size_t compute_base_offset(const TensorShapeVector& shape, const TensorPitches& pitches, int64_t skip_axis) {
|
||||
// in this context, rank can never be < 1, so saving checking overhead
|
||||
auto loop_size = static_cast<int64_t>(shape.size()) - 1;
|
||||
|
||||
size_t base_offset = 0;
|
||||
|
||||
for (int64_t i = 0; i < loop_size; ++i) {
|
||||
if (i != skip_axis)
|
||||
base_offset += shape[i] * pitches[i];
|
||||
}
|
||||
|
||||
return base_offset;
|
||||
}
|
||||
|
||||
// This method computes the number of 'inner_dimension' chunks
|
||||
// Compute the number of 'inner_dimension' elements
|
||||
// Example: input = [2, 3] output = 2
|
||||
// input = [3, 2, 4] output = 3 * 2 = 6
|
||||
// input = [2] output = 1
|
||||
static int64_t calculate_num_inner_dim(const TensorShape& dims) {
|
||||
// in this context, rank can never be < 1, so saving checking overhead
|
||||
static int64_t CalculateInnerDimCount(const TensorShape& dims) {
|
||||
// Rank can never be < 1, no need to check
|
||||
return dims.SizeToDimension(dims.NumDimensions() - 1);
|
||||
}
|
||||
|
||||
// This method computes increments over an 'inner_dimension'
|
||||
// Example 1: current_dims = [0, x] tensor_dims = [3, 1], then current_dims = [1, x]
|
||||
// current_dims = [1, x] tensor_dims = [3, 1], then current_dims = [2, x]
|
||||
// current_dims = [2, x] tensor_dims = [3, 1], then current_dims = [0, x]
|
||||
|
||||
// Example 2: current_dims = [0, 0, x] tensor_dims = [1, 2, 2], then current_dims = [0, 1, x]
|
||||
// current_dims = [0, 1, x] tensor_dims = [1, 2, 2], then current_dims = [0, 0, x]
|
||||
static inline void increment_over_inner_dim(TensorShapeVector& current_dims, const TensorShape& tensor_dims) {
|
||||
// Computes the offset into the input array given the inner_dim count
|
||||
//
|
||||
// If the input indices tensor matched the input tensor this would simply just be inner_dim_size * inner_dim
|
||||
// But since the indices tensor can be smaller we need to do the math based on the smaller size, as the
|
||||
// output tensor size matches the input indices tensor size.
|
||||
//
|
||||
// The calculation is fairly straightforward, starting with the second to innermost axis we muldiv the inner_dim
|
||||
// by the indices shape size. We also skip this calculation on the skip_axis as that's handled elsewhere.
|
||||
//
|
||||
static inline size_t CalculateOffset(size_t inner_dim, const TensorPitches& input_shape_pitches, size_t skip_axis,
|
||||
const TensorShape& indices_shape) {
|
||||
// in this context, rank can never be < 1, so saving checking overhead
|
||||
int64_t rank = static_cast<int64_t>(current_dims.size());
|
||||
|
||||
// 'reset' innermost dimension value
|
||||
current_dims[rank - 1] = 0;
|
||||
size_t rank = input_shape_pitches.size();
|
||||
|
||||
// nothing to increment over
|
||||
if (rank == 1) {
|
||||
return;
|
||||
return 0;
|
||||
}
|
||||
|
||||
int64_t current_axis = rank - 2;
|
||||
size_t base_offset = 0;
|
||||
|
||||
while (current_axis >= 0) {
|
||||
if (++current_dims[current_axis] != tensor_dims[current_axis])
|
||||
return;
|
||||
|
||||
current_dims[current_axis] = 0;
|
||||
--current_axis;
|
||||
for (size_t axis = rank - 1; axis-- > 0;) {
|
||||
auto dim = indices_shape[axis];
|
||||
if (axis != skip_axis)
|
||||
base_offset += (inner_dim % dim) * input_shape_pitches[axis];
|
||||
inner_dim /= dim;
|
||||
}
|
||||
|
||||
return base_offset;
|
||||
}
|
||||
|
||||
#if defined(_MSC_VER)
|
||||
|
|
@ -98,7 +78,7 @@ FORCEINLINE int64_t GetIndex(size_t i, const T* indices, int64_t axis_size) {
|
|||
if (index < 0) // Handle negative indices
|
||||
index += axis_size;
|
||||
if (std::make_unsigned_t<T>(index) >= std::make_unsigned_t<T>(axis_size))
|
||||
ORT_THROW("GatherElements op: Value in indices must be within bounds [-", axis_size, " , ", axis_size - 1, "]. Actual value is ", indices[i]);
|
||||
ORT_THROW("Index out of range");
|
||||
return index;
|
||||
};
|
||||
|
||||
|
|
@ -110,7 +90,8 @@ FORCEINLINE int64_t GetIndex(size_t i, const T* indices, int64_t axis_size) {
|
|||
#endif
|
||||
|
||||
template <typename Tin>
|
||||
static void core_impl(const Tensor* input_tensor, const Tensor* indices_tensor, Tensor* output_tensor, int64_t axis) {
|
||||
static void core_impl(const Tensor* input_tensor, const Tensor* indices_tensor, Tensor* output_tensor, int64_t axis,
|
||||
concurrency::ThreadPool* ttp) {
|
||||
// Get input & output pointers
|
||||
const int8_t* input_data = reinterpret_cast<const int8_t*>(input_tensor->DataRaw());
|
||||
int8_t* output_data = reinterpret_cast<int8_t*>(output_tensor->MutableDataRaw());
|
||||
|
|
@ -120,56 +101,57 @@ static void core_impl(const Tensor* input_tensor, const Tensor* indices_tensor,
|
|||
const int64_t input_rank = static_cast<int64_t>(input_tensor->Shape().NumDimensions());
|
||||
|
||||
const TensorShape& indices_shape = indices_tensor->Shape();
|
||||
size_t num_inner_dim = calculate_num_inner_dim(indices_shape);
|
||||
size_t num_inner_dim = CalculateInnerDimCount(indices_shape);
|
||||
size_t inner_dim_size = indices_shape[input_rank - 1];
|
||||
const Tin* indices = indices_tensor->Data<Tin>();
|
||||
const Tin* indices_data = indices_tensor->Data<Tin>();
|
||||
|
||||
TensorShapeVector process_dims(input_rank, 0);
|
||||
const TensorPitches input_shape_pitches(*input_tensor);
|
||||
int64_t axis_pitch = input_shape_pitches[axis];
|
||||
int64_t axis_size = input_tensor->Shape()[axis];
|
||||
|
||||
auto DoAxis = [](auto* output, auto* input, auto* indices, size_t inner_dim_size, int64_t axis_size, int64_t axis_pitch) {
|
||||
for (size_t i = 0; i < inner_dim_size; i++)
|
||||
output[i] = input[GetIndex(i, indices, axis_size) * axis_pitch + i];
|
||||
};
|
||||
bool innermost_axis = axis == input_rank - 1;
|
||||
bool index_error = false;
|
||||
|
||||
// Special case required for innermost axis, no axis_pitch multiply needed or adding i
|
||||
auto DoInnermostAxis = [](auto* output, auto* input, auto* indices, size_t inner_dim_size, int64_t axis_size, int64_t /*axis_pitch*/) {
|
||||
for (size_t i = 0; i < inner_dim_size; i++)
|
||||
output[i] = input[GetIndex(i, indices, axis_size)];
|
||||
};
|
||||
auto MainLoop = [&](auto* output_data, auto* input_data) {
|
||||
auto BatchWork = [&](size_t inner_dim) {
|
||||
ORT_TRY {
|
||||
auto output = output_data + inner_dim_size * inner_dim;
|
||||
auto input = input_data + CalculateOffset(inner_dim, input_shape_pitches, axis, indices_shape);
|
||||
auto indices = indices_data + inner_dim_size * inner_dim;
|
||||
|
||||
auto MainLoop = [&](auto* output, auto* input, auto LoopFn) {
|
||||
while (num_inner_dim-- != 0) {
|
||||
LoopFn(output, input + compute_base_offset(process_dims, input_shape_pitches, axis), indices, inner_dim_size, axis_size, axis_pitch);
|
||||
output += inner_dim_size;
|
||||
indices += static_cast<Tin>(inner_dim_size);
|
||||
increment_over_inner_dim(process_dims, indices_shape);
|
||||
}
|
||||
if (innermost_axis) {
|
||||
for (size_t i = 0; i < inner_dim_size; i++)
|
||||
output[i] = input[GetIndex(i, indices, axis_size)];
|
||||
} else {
|
||||
for (size_t i = 0; i < inner_dim_size; i++)
|
||||
output[i] = input[GetIndex(i, indices, axis_size) * axis_pitch + i];
|
||||
}
|
||||
}
|
||||
ORT_CATCH(const std::exception&) {
|
||||
index_error = true;
|
||||
}
|
||||
};
|
||||
|
||||
concurrency::ThreadPool::TryBatchParallelFor(ttp, num_inner_dim, BatchWork, 0);
|
||||
};
|
||||
|
||||
// Iterate over the elements based on the element size (or if it's a string). For everything but strings
|
||||
// we do a binary copy, so handling the 1, 2, 4, and 8 byte sizes covers all cases
|
||||
auto SwitchOnSizes = [&](auto LoopFn) {
|
||||
if (is_string)
|
||||
MainLoop(reinterpret_cast<std::string*>(output_data), reinterpret_cast<const std::string*>(input_data), LoopFn);
|
||||
else if (element_size == sizeof(uint32_t))
|
||||
MainLoop(reinterpret_cast<uint32_t*>(output_data), reinterpret_cast<const uint32_t*>(input_data), LoopFn);
|
||||
else if (element_size == sizeof(uint16_t))
|
||||
MainLoop(reinterpret_cast<uint16_t*>(output_data), reinterpret_cast<const uint16_t*>(input_data), LoopFn);
|
||||
else if (element_size == sizeof(uint8_t))
|
||||
MainLoop(reinterpret_cast<uint8_t*>(output_data), reinterpret_cast<const uint8_t*>(input_data), LoopFn);
|
||||
else if (element_size == sizeof(uint64_t))
|
||||
MainLoop(reinterpret_cast<uint64_t*>(output_data), reinterpret_cast<const uint64_t*>(input_data), LoopFn);
|
||||
else
|
||||
ORT_THROW("GatherElements op: Unsupported tensor type, size:", element_size);
|
||||
};
|
||||
|
||||
if (axis == input_rank - 1)
|
||||
SwitchOnSizes(DoInnermostAxis);
|
||||
if (is_string)
|
||||
MainLoop(reinterpret_cast<std::string*>(output_data), reinterpret_cast<const std::string*>(input_data));
|
||||
else if (element_size == sizeof(uint32_t))
|
||||
MainLoop(reinterpret_cast<uint32_t*>(output_data), reinterpret_cast<const uint32_t*>(input_data));
|
||||
else if (element_size == sizeof(uint16_t))
|
||||
MainLoop(reinterpret_cast<uint16_t*>(output_data), reinterpret_cast<const uint16_t*>(input_data));
|
||||
else if (element_size == sizeof(uint8_t))
|
||||
MainLoop(reinterpret_cast<uint8_t*>(output_data), reinterpret_cast<const uint8_t*>(input_data));
|
||||
else if (element_size == sizeof(uint64_t))
|
||||
MainLoop(reinterpret_cast<uint64_t*>(output_data), reinterpret_cast<const uint64_t*>(input_data));
|
||||
else
|
||||
SwitchOnSizes(DoAxis);
|
||||
ORT_THROW("GatherElements op: Unsupported tensor type, size:", element_size);
|
||||
|
||||
if (index_error)
|
||||
ORT_THROW("GatherElements op: Out of range value in index tensor");
|
||||
}
|
||||
#ifdef __GNUC__
|
||||
#pragma GCC diagnostic pop
|
||||
|
|
@ -233,10 +215,11 @@ Status GatherElements::Compute(OpKernelContext* context) const {
|
|||
if (indices_shape.Size() == 0)
|
||||
return Status::OK();
|
||||
|
||||
auto* ttp = context->GetOperatorThreadPool();
|
||||
if (indices_tensor->IsDataType<int32_t>())
|
||||
core_impl<int32_t>(input_tensor, indices_tensor, output_tensor, axis);
|
||||
core_impl<int32_t>(input_tensor, indices_tensor, output_tensor, axis, ttp);
|
||||
else
|
||||
core_impl<int64_t>(input_tensor, indices_tensor, output_tensor, axis);
|
||||
core_impl<int64_t>(input_tensor, indices_tensor, output_tensor, axis, ttp);
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -95,7 +95,7 @@ void RunTypedTest() {
|
|||
// skip openvino which will not throw error message but will ensure no out-of-bound access
|
||||
// skip TensorRT because it doesn't support out of bounds indices
|
||||
test5.Run(OpTester::ExpectResult::kExpectFailure,
|
||||
"GatherElements op: Value in indices must be within bounds [-2 , 1]. Actual value is 2",
|
||||
"GatherElements op: Out of range value in index tensor",
|
||||
{kNupharExecutionProvider, kCudaExecutionProvider, kRocmExecutionProvider, kOpenVINOExecutionProvider, kTensorrtExecutionProvider});
|
||||
|
||||
// 3D input - axis 1
|
||||
|
|
@ -250,7 +250,7 @@ void RunTypedTest<std::string>() {
|
|||
// skip nuphar, which will not throw error message but will ensure no out-of-bound access
|
||||
// skip Openvino, which will not throw error message but will ensure no out-of-bound access
|
||||
test4.Run(OpTester::ExpectResult::kExpectFailure,
|
||||
"GatherElements op: Value in indices must be within bounds [-2 , 1]. Actual value is -3",
|
||||
"GatherElements op: Out of range value in index tensor",
|
||||
{kNupharExecutionProvider, kOpenVINOExecutionProvider});
|
||||
|
||||
// 3D input - axis 1
|
||||
|
|
|
|||
Loading…
Reference in a new issue