Rework broadcasting setup to decrease binary size. (#5227)

* Rework broadcasting setup to decrease binary size. Push all the type specific down and separate out the broadcasting/parallelization.

Reductions:
element_wise_ops: 521.0KB -> 268.8KB
where: 25.8 KB -> 17.3 KB
qlinear_binary_op: 28.1 -> 12.8
This commit is contained in:
Scott McKay 2020-09-23 14:15:40 +10:00 committed by GitHub
parent 43faf9e388
commit c52561d044
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 1161 additions and 706 deletions

View file

@ -12,52 +12,45 @@ using onnxruntime::concurrency::ThreadPool;
namespace onnxruntime {
namespace contrib {
template <typename T, typename Input0Scalar, typename Input1Scalar, typename General>
void QLinearBroadcastLoop(TBroadcaster<T, T>& bc, TBroadcastOutput<T>& output, Input0Scalar input0scalar, Input1Scalar input1scalar, General general,
float A_scale, float B_scale, float C_scale, T A_zero_point, T B_zero_point, T C_zero_point) {
if (bc.IsInput0Scalar()) {
while (output)
input0scalar(output.NextSpanOutput(), bc.NextScalar0(), bc.NextSpan1(), A_scale, B_scale, C_scale, A_zero_point, B_zero_point, C_zero_point);
} else if (bc.IsInput1Scalar()) {
while (output)
input1scalar(output.NextSpanOutput(), bc.NextSpan0(), bc.NextScalar1(), A_scale, B_scale, C_scale, A_zero_point, B_zero_point, C_zero_point);
} else {
while (output)
general(output.NextSpanOutput(), bc.NextSpan0(), bc.NextSpan1(), A_scale, B_scale, C_scale, A_zero_point, B_zero_point, C_zero_point);
namespace {
struct QLinearBroadcastHelper : public BroadcastHelper {
QLinearBroadcastHelper(InputBroadcaster& input_broadcaster,
OutputBroadcaster& output_broadcaster,
ThreadPool* threadpool,
double unit_cost,
float A_scale_in, float B_scale_in, float C_scale_in,
uint8_t A_zero_point_in, uint8_t B_zero_point_in, uint8_t C_zero_point_in)
: BroadcastHelper{input_broadcaster, output_broadcaster, nullptr, threadpool, unit_cost},
A_scale{A_scale_in},
B_scale{B_scale_in},
C_scale{C_scale_in},
A_zero_point{A_zero_point_in},
B_zero_point{B_zero_point_in},
C_zero_point{C_zero_point_in} {
}
}
template <typename T, typename Input0Scalar, typename Input1Scalar, typename General>
void QLinearBroadcastOneSpan(ThreadPool* tp, double unit_cost,
gsl::span<T> output_span, gsl::span<const T> input0_span, gsl::span<const T> input1_span,
Input0Scalar input0scalar, Input1Scalar input1scalar, General general,
float A_scale, float B_scale, float C_scale, T A_zero_point, T B_zero_point, T C_zero_point) {
if (input0_span.size() == 1) {
ThreadPool::TryParallelFor(tp, output_span.size(), unit_cost,
[=](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
input0scalar(output_span.subspan(first, count), *input0_span.data(), input1_span.subspan(first, count),
A_scale, B_scale, C_scale, A_zero_point, B_zero_point, C_zero_point);
});
} else if (input1_span.size() == 1) {
ThreadPool::TryParallelFor(tp, output_span.size(), unit_cost,
[=](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
input1scalar(output_span.subspan(first, count), input0_span.subspan(first, count), *input1_span.data(),
A_scale, B_scale, C_scale, A_zero_point, B_zero_point, C_zero_point);
});
} else {
ThreadPool::TryParallelFor(tp, output_span.size(), unit_cost,
[=](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
general(output_span.subspan(first, count), input0_span.subspan(first, count), input1_span.subspan(first, count),
A_scale, B_scale, C_scale, A_zero_point, B_zero_point, C_zero_point);
});
QLinearBroadcastHelper(const QLinearBroadcastHelper& rhs, size_t offset, size_t num_elements)
: BroadcastHelper(rhs, offset, num_elements),
A_scale{rhs.A_scale},
B_scale{rhs.B_scale},
C_scale{rhs.C_scale},
A_zero_point{rhs.A_zero_point},
B_zero_point{rhs.B_zero_point},
C_zero_point{rhs.C_zero_point} {
}
}
template <typename T, typename Input0Scalar, typename Input1Scalar, typename General>
Status QLinearBroadcastTwo(OpKernelContext& context, Input0Scalar input0scalar, Input1Scalar input1scalar, General general, double unit_cost) {
float A_scale;
float B_scale;
float C_scale;
// storage for these is uint8_t but original value may be uint8_t or int8_t.
// typed code that uses values needs to cast to correct representation
uint8_t A_zero_point;
uint8_t B_zero_point;
uint8_t C_zero_point;
};
template <typename T>
void QLinearImpl(OpKernelContext& context, double unit_cost, const ProcessBroadcastSpanFuncs& functors) {
auto tensor_a_scale = context.Input<Tensor>(1);
auto tensor_a_zero_point = context.Input<Tensor>(2);
auto tensor_b_scale = context.Input<Tensor>(4);
@ -79,87 +72,119 @@ Status QLinearBroadcastTwo(OpKernelContext& context, Input0Scalar input0scalar,
"MatmulInteger : input1 C_zero_point must be a scalar or 1D tensor of size 1 if given");
const float A_scale = *(tensor_a_scale->Data<float>());
const T A_zero_point = (nullptr == tensor_a_zero_point) ? static_cast<T>(0) : *(tensor_a_zero_point->template Data<T>());
const T A_zero_point = (nullptr == tensor_a_zero_point) ? T{} : *(tensor_a_zero_point->template Data<T>());
const float B_scale = *(tensor_b_scale->Data<float>());
const T B_zero_point = (nullptr == tensor_b_zero_point) ? static_cast<T>(0) : *(tensor_b_zero_point->template Data<T>());
const T B_zero_point = (nullptr == tensor_b_zero_point) ? T{} : *(tensor_b_zero_point->template Data<T>());
const float C_scale = *(tensor_c_scale->Data<float>());
const T C_zero_point = (nullptr == tensor_c_zero_point) ? static_cast<T>(0) : *(tensor_c_zero_point->template Data<T>());
const T C_zero_point = (nullptr == tensor_c_zero_point) ? T{} : *(tensor_c_zero_point->template Data<T>());
TBroadcaster<T, T> bc(*context.Input<Tensor>(0), *context.Input<Tensor>(3));
Tensor& output_tensor = *context.Output(0, bc.GetOutputShape());
auto span_size = bc.GetSpanSize();
TBroadcastOutput<T> output(span_size, output_tensor);
int64_t output_len = output_tensor.Shape().Size();
InputBroadcaster input_broadcaster{*context.Input<Tensor>(0), *context.Input<Tensor>(3)};
OutputBroadcaster output_broadcaster{input_broadcaster.GetSpanSize(),
*context.Output(0, input_broadcaster.GetOutputShape())};
QLinearBroadcastHelper broadcast_helper(input_broadcaster, output_broadcaster,
context.GetOperatorThreadPool(), unit_cost,
A_scale, B_scale, C_scale,
static_cast<uint8_t>(A_zero_point),
static_cast<uint8_t>(B_zero_point),
static_cast<uint8_t>(C_zero_point));
BroadcastLooper(broadcast_helper, functors);
}
} // namespace
template <typename T>
Status QLinearAdd<T>::Compute(OpKernelContext* context) const {
const ProcessBroadcastSpanFuncs functors = {
[](BroadcastHelper& per_iter_bh) {
QLinearBroadcastHelper& qlbh = static_cast<QLinearBroadcastHelper&>(per_iter_bh);
const T input0 = per_iter_bh.ScalarInput0<T>();
auto input1 = per_iter_bh.SpanInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
MlasQLinearAdd(input1.data(),
qlbh.B_scale, static_cast<T>(qlbh.B_zero_point),
&input0,
qlbh.A_scale, static_cast<T>(qlbh.A_zero_point),
qlbh.C_scale, static_cast<T>(qlbh.C_zero_point),
output.data(), output.size(), true);
},
[](BroadcastHelper& per_iter_bh) {
QLinearBroadcastHelper& qlbh = static_cast<QLinearBroadcastHelper&>(per_iter_bh);
auto input0 = per_iter_bh.SpanInput0<T>();
const T input1 = per_iter_bh.ScalarInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
MlasQLinearAdd(input0.data(),
qlbh.A_scale, static_cast<T>(qlbh.A_zero_point),
&input1,
qlbh.B_scale, static_cast<T>(qlbh.B_zero_point),
qlbh.C_scale, static_cast<T>(qlbh.C_zero_point),
output.data(), output.size(), true);
},
[](BroadcastHelper& per_iter_bh) {
QLinearBroadcastHelper& qlbh = static_cast<QLinearBroadcastHelper&>(per_iter_bh);
auto input0 = per_iter_bh.SpanInput0<T>();
auto input1 = per_iter_bh.SpanInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
MlasQLinearAdd(input0.data(),
qlbh.A_scale, static_cast<T>(qlbh.A_zero_point),
input1.data(),
qlbh.B_scale, static_cast<T>(qlbh.B_zero_point),
qlbh.C_scale, static_cast<T>(qlbh.C_zero_point),
output.data(), output.size(), false);
}};
QLinearImpl<T>(*context, 1.0, functors);
ThreadPool* tp = context.GetOperatorThreadPool();
if (output_len == static_cast<int64_t>(span_size)) { // Only one big span for all data, parallel inside it
auto span0 = bc.IsInput0Scalar() ? gsl::span<const T>(&bc.NextScalar0(), 1) : bc.NextSpan0();
auto span1 = bc.IsInput1Scalar() ? gsl::span<const T>(&bc.NextScalar1(), 1) : bc.NextSpan1();
QLinearBroadcastOneSpan(tp, unit_cost, output.NextSpanOutput(), span0, span1,
input0scalar, input1scalar, general,
A_scale, B_scale, C_scale, A_zero_point, B_zero_point, C_zero_point);
} else {
ThreadPool::TryParallelFor(
tp, output_len / span_size, unit_cost * span_size,
[=, &bc, &output_tensor](std::ptrdiff_t first_span, std::ptrdiff_t last_span) {
TBroadcaster<T, T> span_bc(bc);
TBroadcastOutput<T> span_output(span_size, output_tensor, first_span * span_size, last_span * span_size);
span_bc.AdvanceBy(first_span * span_size);
QLinearBroadcastLoop(span_bc, span_output, input0scalar, input1scalar, general,
A_scale, B_scale, C_scale, A_zero_point, B_zero_point, C_zero_point);
});
}
return Status::OK();
}
template <typename T>
Status QLinearAdd<T>::Compute(OpKernelContext* context) const {
return QLinearBroadcastTwo<T>(
*context,
[](gsl::span<T> output, const T& input0, gsl::span<const T> input1,
float A_scale, float B_scale, float C_scale, T A_zero_point, T B_zero_point, T C_zero_point) {
MlasQLinearAdd(input1.data(), B_scale, B_zero_point,
&input0, A_scale, A_zero_point,
C_scale, C_zero_point, output.data(), output.size(), true);
},
[](gsl::span<T> output, gsl::span<const T> input0, const T& input1,
float A_scale, float B_scale, float C_scale, T A_zero_point, T B_zero_point, T C_zero_point) {
MlasQLinearAdd(input0.data(), A_scale, A_zero_point,
&input1, B_scale, B_zero_point,
C_scale, C_zero_point, output.data(), output.size(), true);
},
[](gsl::span<T> output, gsl::span<const T> input0, gsl::span<const T> input1,
float A_scale, float B_scale, float C_scale, T A_zero_point, T B_zero_point, T C_zero_point) {
MlasQLinearAdd(input0.data(), A_scale, A_zero_point,
input1.data(), B_scale, B_zero_point,
C_scale, C_zero_point, output.data(), output.size(), false);
},
1.0);
}
template <typename T>
Status QLinearMul<T>::Compute(OpKernelContext* context) const {
return QLinearBroadcastTwo<T>(
*context,
[](gsl::span<T> output, const T& input0, gsl::span<const T> input1,
float A_scale, float B_scale, float C_scale, T A_zero_point, T B_zero_point, T C_zero_point) {
MlasQLinearMul(input1.data(), B_scale, B_zero_point,
&input0, A_scale, A_zero_point,
C_scale, C_zero_point, output.data(), output.size(), true);
const ProcessBroadcastSpanFuncs functors = {
[](BroadcastHelper& per_iter_bh) {
QLinearBroadcastHelper& qlbh = static_cast<QLinearBroadcastHelper&>(per_iter_bh);
const T input0 = per_iter_bh.ScalarInput0<T>();
auto input1 = per_iter_bh.SpanInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
MlasQLinearMul(input1.data(),
qlbh.B_scale, static_cast<T>(qlbh.B_zero_point),
&input0,
qlbh.A_scale, static_cast<T>(qlbh.A_zero_point),
qlbh.C_scale, static_cast<T>(qlbh.C_zero_point),
output.data(), output.size(), true);
},
[](gsl::span<T> output, gsl::span<const T> input0, const T& input1,
float A_scale, float B_scale, float C_scale, T A_zero_point, T B_zero_point, T C_zero_point) {
MlasQLinearMul(input0.data(), A_scale, A_zero_point,
&input1, B_scale, B_zero_point,
C_scale, C_zero_point, output.data(), output.size(), true);
[](BroadcastHelper& per_iter_bh) {
QLinearBroadcastHelper& qlbh = static_cast<QLinearBroadcastHelper&>(per_iter_bh);
auto input0 = per_iter_bh.SpanInput0<T>();
const T input1 = per_iter_bh.ScalarInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
MlasQLinearMul(input0.data(),
qlbh.A_scale, static_cast<T>(qlbh.A_zero_point),
&input1,
qlbh.B_scale, static_cast<T>(qlbh.B_zero_point),
qlbh.C_scale, static_cast<T>(qlbh.C_zero_point),
output.data(), output.size(), true);
},
[](gsl::span<T> output, gsl::span<const T> input0, gsl::span<const T> input1,
float A_scale, float B_scale, float C_scale, T A_zero_point, T B_zero_point, T C_zero_point) {
MlasQLinearMul(input0.data(), A_scale, A_zero_point,
input1.data(), B_scale, B_zero_point,
C_scale, C_zero_point, output.data(), output.size(), false);
},
1.0);
[](BroadcastHelper& per_iter_bh) {
QLinearBroadcastHelper& qlbh = static_cast<QLinearBroadcastHelper&>(per_iter_bh);
auto input0 = per_iter_bh.SpanInput0<T>();
auto input1 = per_iter_bh.SpanInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
MlasQLinearMul(input0.data(),
qlbh.A_scale, static_cast<T>(qlbh.A_zero_point),
input1.data(),
qlbh.B_scale, static_cast<T>(qlbh.B_zero_point),
qlbh.C_scale, static_cast<T>(qlbh.C_zero_point),
output.data(), output.size(), false);
}};
QLinearImpl<T>(*context, 1.0, functors);
return Status::OK();
}
#define REG_QLINEAR_ELEMENTWISE_TYPED_KERNEL(op_name, version, data_type, KERNEL_CLASS) \

File diff suppressed because it is too large Load diff

View file

@ -11,7 +11,7 @@
namespace onnxruntime {
namespace functors {
template<typename T>
template <typename T>
struct Log final : public ElementWiseRangedTransform<T> {
Status Init(const onnxruntime::NodeAttributes) {
return Status::OK();
@ -22,13 +22,13 @@ struct Log final : public ElementWiseRangedTransform<T> {
using T2 = typename std::remove_const<T1>::type;
return new T2(*this);
}
float Cost() const final { return 15.0f; }
void operator()(std::ptrdiff_t first, std::ptrdiff_t last) const {
ptrdiff_t len = last - first;
T* output_ptr = this->output + first;
ConstEigenVectorArrayMap<T> xm(this->input + first, len);
T* output_ptr = this->output + first;
ConstEigenVectorArrayMap<T> xm(this->input + first, len);
EigenVectorArrayMap<T> ym(output_ptr, len);
ym = xm.log();
}
@ -194,7 +194,7 @@ struct Exp final : public ElementWiseRangedTransform<T> {
ym = xm.exp();
}
};
}
} // namespace functors
DEFINE_ELE_KERNEL(Log)
DEFINE_ELE_KERNEL(Abs)
@ -205,7 +205,6 @@ DEFINE_ELE_KERNEL(Reciprocal)
DEFINE_ELE_KERNEL(Sqrt)
DEFINE_ELE_KERNEL(Exp)
template <typename T>
class Add final : public OpKernel {
public:
@ -242,7 +241,6 @@ class Div final : public OpKernel {
Status Compute(OpKernelContext* context) const override;
};
class Pow final : public OpKernel {
public:
Pow(const OpKernelInfo& info) : OpKernel(info) {
@ -312,10 +310,9 @@ class Max_8 final : public OpKernel {
Status Compute(OpKernelContext* context) const override;
private:
template<typename T>
struct ComputeImpl;
private:
template <typename T>
struct ComputeImpl;
};
class Not final : public OpKernel {
@ -434,11 +431,18 @@ class Erf final : public OpKernel {
};
template <typename T>
auto MakeEigenArrayMap(Tensor& t) -> EigenVectorArrayMap<T> { return EigenVectorArrayMap<T>(t.template MutableData<T>(), t.Shape().Size()); }
auto MakeEigenArrayMap(Tensor& t) -> EigenVectorArrayMap<T> {
return EigenVectorArrayMap<T>(t.template MutableData<T>(), t.Shape().Size());
}
template <typename T>
auto MakeEigenArrayMap(const Tensor& t) -> ConstEigenVectorArrayMap<T> { return ConstEigenVectorArrayMap<T>(t.template Data<T>(), t.Shape().Size()); }
auto MakeEigenArrayMap(const Tensor& t) -> ConstEigenVectorArrayMap<T> {
return ConstEigenVectorArrayMap<T>(t.template Data<T>(), t.Shape().Size());
}
struct BroadcastIterator {
size_t Current() const { return index_; }
size_t AdvanceBy(size_t delta) {
size_t index = index_;
@ -452,7 +456,7 @@ struct BroadcastIterator {
break;
counters_[counterIndex] = 0;
}
} else if (counters_[0] > counts_[0]) { // Keep original logic above so that in most case it is faster
} else if (counters_[0] > counts_[0]) { // Keep original logic above so that in most case it is faster
delta = counters_[0] / counts_[0];
counters_[0] = counters_[0] % counts_[0];
for (size_t counterIndex = 1; counterIndex < counters_.size(); counterIndex++) {
@ -623,15 +627,21 @@ struct Broadcaster {
std::vector<int64_t> output_shape_;
};
template <typename T0, typename T1>
struct TBroadcaster {
TBroadcaster(const Tensor& input0, const Tensor& input1)
struct InputBroadcaster {
InputBroadcaster(const Tensor& input0, const Tensor& input1)
: input_tensor0_(input0),
input_tensor1_(input1) {
input_tensor1_(&input1),
input_tensor1_shape_(input1.Shape()) {
}
InputBroadcaster(const Tensor& input0, const TensorShape& input1_shape)
: input_tensor0_(input0),
input_tensor1_(nullptr),
input_tensor1_shape_(input1_shape) {
}
void AdvanceBy(size_t offset) {
ORT_ENFORCE(offset % span_size_ == 0, "TBroadcaster can only start at span boundary!");
ORT_ENFORCE(offset % span_size_ == 0, "InputBroadcaster can only start at span boundary!");
broadcaster_.iterator1_.AdvanceBy(offset);
broadcaster_.iterator2_.AdvanceBy(offset);
}
@ -639,80 +649,312 @@ struct TBroadcaster {
TensorShape GetOutputShape() const { return TensorShape(broadcaster_.output_shape_); }
size_t GetSpanSize() const { return span_size_; }
// Check whether we have a tensor instance for input 1. Code using this class is required to validate this
// before calling any methods that require input 1 to have data.
bool HaveTwoTensors() const { return input_tensor1_ != nullptr; }
bool IsInput0Scalar() const { return broadcaster_.iterator1_.deltas_.front() == 0; }
bool IsInput1Scalar() const { return broadcaster_.iterator2_.deltas_.front() == 0; }
const T0& NextScalar0() { return *Next0(); }
const T1& NextScalar1() { return *Next1(); }
size_t Input0ElementSize() const { return input0_element_size_; }
size_t Input1ElementSize() const { return input1_element_size_; }
gsl::span<const T0> NextSpan0() { return gsl::span<const T0>(Next0(), span_size_); }
gsl::span<const T1> NextSpan1() { return gsl::span<const T1>(Next1(), span_size_); }
template <typename T>
const T& Scalar0() { return *(static_cast<const T*>(input0_bytes_) + broadcaster_.iterator1_.Current()); }
template <typename T>
const T& Scalar1() { return *(static_cast<const T*>(input1_bytes_) + broadcaster_.iterator2_.Current()); }
ConstEigenVectorMap<T0> NextEigen0() { return ConstEigenVectorMap<T0>(Next0(), span_size_); }
ConstEigenVectorMap<T1> NextEigen1() { return ConstEigenVectorMap<T1>(Next1(), span_size_); }
// general usage is to get a full span, but if we parallelize within a span we need intra-span pieces
// which are specified via offset and num_elements
template <typename T>
ConstEigenVectorMap<T> Eigen0(size_t offset, size_t num_elements) {
assert(offset < span_size_ && (offset + num_elements) <= span_size_);
return ConstEigenVectorMap<T>(static_cast<const T*>(input0_bytes_) + broadcaster_.iterator1_.Current() + offset,
num_elements);
}
template <typename T>
ConstEigenVectorMap<T> Eigen1(size_t offset, size_t num_elements) {
assert(offset < span_size_ && (offset + num_elements) <= span_size_);
return ConstEigenVectorMap<T>(static_cast<const T*>(input1_bytes_) + broadcaster_.iterator2_.Current() + offset,
num_elements);
}
template <typename T>
gsl::span<const T> Span0(size_t offset, size_t num_elements) {
return gsl::span<const T>(static_cast<const T*>(input0_bytes_) + broadcaster_.iterator1_.Current() + offset,
num_elements);
}
template <typename T>
gsl::span<const T> Span1(size_t offset, size_t num_elements) {
return gsl::span<const T>(static_cast<const T*>(input1_bytes_) + broadcaster_.iterator2_.Current() + offset,
num_elements);
}
void Next() {
AdvanceBy(span_size_);
}
private:
const T0* Next0() { return input0_ + broadcaster_.iterator1_.AdvanceBy(span_size_); }
const T1* Next1() { return input1_ + broadcaster_.iterator2_.AdvanceBy(span_size_); }
const Tensor& input_tensor0_;
const Tensor& input_tensor1_;
Broadcaster broadcaster_{input_tensor0_.Shape().GetDims(), input_tensor1_.Shape().GetDims()};
size_t span_size_{broadcaster_.GetSpanSize()};
// need to support use case where input1 is just the shape for Expand op
const Tensor* input_tensor1_{nullptr};
const TensorShape& input_tensor1_shape_;
const size_t input0_element_size_{input_tensor0_.DataType()->Size()};
const size_t input1_element_size_{input_tensor1_ ? input_tensor1_->DataType()->Size() : 0};
const void* input0_bytes_{input_tensor0_.DataRaw()};
const void* input1_bytes_{input_tensor1_ ? input_tensor1_->DataRaw() : nullptr};
const T0* input0_{input_tensor0_.template Data<T0>()};
const T1* input1_{input_tensor1_.template Data<T1>()};
Broadcaster broadcaster_{input_tensor0_.Shape().GetDims(), input_tensor1_shape_.GetDims()};
size_t span_size_{broadcaster_.GetSpanSize()};
};
template <typename T>
struct TBroadcastOutput {
TBroadcastOutput(size_t span_size, Tensor& tensor, int64_t start_offset = 0, int64_t end_offset = 0)
: span_size_(span_size) {
struct OutputBroadcaster {
OutputBroadcaster(size_t span_size, Tensor& tensor, int64_t start_offset = 0, int64_t end_offset = 0)
: element_size_(tensor.DataType()->Size()),
span_size_(span_size) {
int64_t len = tensor.Shape().Size();
int64_t real_end = (end_offset <= 0) ? len : end_offset;
if (start_offset != 0 || end_offset != 0) { // Keep original semantic
if (start_offset != 0 || end_offset != 0) { // Keep original semantic
ORT_ENFORCE(start_offset >= 0 && real_end >= 0 && start_offset <= real_end && real_end <= len,
"Invalid start/ending offset [", start_offset, ",", real_end, ") for tensor of length:", len);
ORT_ENFORCE(start_offset % span_size == 0 && real_end % span_size == 0,
"Broadcast Output range [", start_offset, ", ", real_end,
") are not at boundary of span with size:", span_size);
}
output_ = tensor.template MutableData<T>() + start_offset;
output_end_ = tensor.template MutableData<T>() + real_end;
output_elements_ = real_end - start_offset;
output_bytes_ = static_cast<uint8_t*>(tensor.MutableDataRaw()) + (start_offset * element_size_);
output_end_ = output_bytes_ + ((real_end - start_offset) * element_size_);
}
size_t OutputElementSize() const { return element_size_; }
size_t NumOutputElements() const { return output_elements_; }
operator bool() const {
return output_ != output_end_;
return output_bytes_ != output_end_;
}
EigenVectorMap<T> NextEigenOutput() {
return EigenVectorMap<T>(NextOutput(), span_size_);
template <typename T>
EigenVectorMap<T> EigenOutput(size_t offset, size_t num_elements) {
assert(offset < span_size_ && (offset + num_elements) <= span_size_);
return EigenVectorMap<T>(reinterpret_cast<T*>(output_bytes_) + offset, num_elements);
}
gsl::span<T> NextSpanOutput() {
return gsl::span<T>(NextOutput(), span_size_);
template <typename T>
gsl::span<T> SpanOutput(size_t offset, size_t num_elements) {
assert(offset < span_size_ && (offset + num_elements) <= span_size_);
return gsl::span<T>(reinterpret_cast<T*>(output_bytes_) + offset, num_elements);
}
void Next() {
output_bytes_ += (span_size_ * element_size_);
}
private:
T* NextOutput() {
T* output = output_;
output_ += span_size_;
return output;
}
T* output_;
const T* output_end_;
size_t span_size_;
const size_t element_size_;
const size_t span_size_;
size_t output_elements_;
uint8_t* output_bytes_;
const void* output_end_;
};
template <typename T>
class BroadcastHelper {
public:
// general purpose ctor
BroadcastHelper(InputBroadcaster& input_broadcaster,
OutputBroadcaster& output_broadcaster,
void* user_data = nullptr,
concurrency::ThreadPool* tp = nullptr,
double unit_cost = 0.0)
: input_broadcaster_(input_broadcaster),
output_broadcaster_(output_broadcaster),
threadpool_(tp),
unit_cost_(unit_cost),
user_data_(user_data) {
}
// ctor for use when we parallelize within a span.
BroadcastHelper(const BroadcastHelper& rhs, size_t offset, size_t num_elements)
: input_broadcaster_(rhs.input_broadcaster_),
output_broadcaster_(rhs.output_broadcaster_),
input0_offset_(IsInput0Scalar() ? 0 : offset),
input0_num_elements_(IsInput0Scalar() ? 1 : num_elements),
input1_offset_(IsInput1Scalar() ? 0 : offset),
input1_num_elements_(IsInput1Scalar() ? 1 : num_elements),
output_offset_(offset),
output_num_elements_(num_elements),
user_data_(rhs.user_data_) {
}
// convenience accessors to simplify usage of this class. these will be optimized away in a release build
bool HaveTwoTensorInputs() const { return input_broadcaster_.HaveTwoTensors(); }
bool IsInput0Scalar() const { return input_broadcaster_.IsInput0Scalar(); }
bool IsInput1Scalar() const { return input_broadcaster_.IsInput1Scalar(); }
size_t Input0ElementSize() const { return input_broadcaster_.Input0ElementSize(); }
size_t Input1ElementSize() const { return input_broadcaster_.Input1ElementSize(); }
size_t OutputElementSize() const { return output_broadcaster_.OutputElementSize(); }
size_t NumOutputElements() const { return output_broadcaster_.NumOutputElements(); }
bool SingleSpanOutput() const { return input_broadcaster_.GetSpanSize() == output_broadcaster_.NumOutputElements(); }
template <typename T>
const T& ScalarInput0() { return input_broadcaster_.Scalar0<T>(); }
template <typename T>
const T& ScalarInput1() { return input_broadcaster_.Scalar1<T>(); }
template <typename T>
ConstEigenVectorMap<T> EigenInput0() { return input_broadcaster_.Eigen0<T>(input0_offset_, input0_num_elements_); }
template <typename T>
ConstEigenVectorMap<T> EigenInput1() { return input_broadcaster_.Eigen1<T>(input1_offset_, input1_num_elements_); }
template <typename T>
EigenVectorMap<T> OutputEigen() { return output_broadcaster_.EigenOutput<T>(output_offset_, output_num_elements_); }
template <typename T>
gsl::span<const T> SpanInput0() { return input_broadcaster_.Span0<T>(input0_offset_, input0_num_elements_); }
template <typename T>
gsl::span<const T> SpanInput1() { return input_broadcaster_.Span1<T>(input1_offset_, input1_num_elements_); }
template <typename T>
gsl::span<T> OutputSpan() { return output_broadcaster_.SpanOutput<T>(output_offset_, output_num_elements_); }
void Next() {
input_broadcaster_.Next();
output_broadcaster_.Next();
}
bool NeedMoreOutput() const { return output_broadcaster_; }
concurrency::ThreadPool* Threadpool() const { return threadpool_; }
double UnitCost() const { return unit_cost_; }
// user data is an opaque blob. there is no memory management provided by BroadcastHelper.
// if the BroadcastHelper instance is copied during parallelization the pointer will be copied across
void SetUserData(void* user_data) { user_data_ = user_data; }
void* GetUserData() const { return user_data_; }
private:
InputBroadcaster& input_broadcaster_;
OutputBroadcaster& output_broadcaster_;
// info required if we parallelize within a span
concurrency::ThreadPool* threadpool_{nullptr};
double unit_cost_{0.0};
size_t input0_offset_{0};
size_t input0_num_elements_{input_broadcaster_.GetSpanSize()}; // default all to getting one full span
size_t input1_offset_{0};
size_t input1_num_elements_{input_broadcaster_.GetSpanSize()};
size_t output_offset_{0};
size_t output_num_elements_{input_broadcaster_.GetSpanSize()};
// opaque user data that is passed through
void* user_data_{nullptr};
};
// type agnostic functions to use in the low level broadcasting to process each span.
// type specific logic is applied within the functions.
// Raw function pointer is significantly cheaper in terms of binary size at the cost of no support for captures.
using ProcessSpanFunc = void (*)(BroadcastHelper&);
struct ProcessBroadcastSpanFuncs {
ProcessSpanFunc input0scalar;
ProcessSpanFunc input1scalar;
ProcessSpanFunc general;
};
// Parallelize processing of data where all the output is covered by a single span
template <typename TBroadcastHelper>
static void ParallelizeSingleSpan(TBroadcastHelper& helper, const ProcessBroadcastSpanFuncs& functors) {
TensorOpCost cost{static_cast<float>(std::max(helper.Input0ElementSize(), helper.Input1ElementSize())),
static_cast<float>(helper.OutputElementSize()),
helper.UnitCost()};
if (helper.IsInput0Scalar()) {
concurrency::ThreadPool::TryParallelFor(
helper.Threadpool(), helper.NumOutputElements(), cost,
[&helper, &functors](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
TBroadcastHelper segment_helper(helper, first, count);
functors.input0scalar(segment_helper);
});
} else if (helper.IsInput1Scalar()) {
concurrency::ThreadPool::TryParallelFor(
helper.Threadpool(), helper.NumOutputElements(), cost,
[&helper, &functors](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
TBroadcastHelper segment_helper(helper, first, count);
functors.input1scalar(segment_helper);
});
} else {
concurrency::ThreadPool::TryParallelFor(
helper.Threadpool(), helper.NumOutputElements(), cost,
[&helper, &functors](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
TBroadcastHelper segment_helper(helper, first, count);
functors.general(segment_helper);
});
}
}
// Broadcast two inputs with no parallelization.
//
// This function is type agnostic, and uses function pointers instead of std::function, to minimize binary size.
// Type specific logic is plugged in via the functions in ProcessBroadcastSpanFuncs.
// Optional user_data can be provided, and will be available to the ProcessSpanFunc implementations
// via BroadcastHelper.GetUserData().
void UntypedBroadcastTwo(OpKernelContext& context, const ProcessBroadcastSpanFuncs& funcs, void* user_data = nullptr);
// Broadcast two inputs with parallelization.
//
// Operator usage is the same as the parallelization is opaque to the operator.
// unit_cost must be a valid cost value.
void UntypedBroadcastTwo(OpKernelContext& context, const ProcessBroadcastSpanFuncs& funcs, double unit_cost,
void* user_data = nullptr);
// Helper to provide the looping logic with optimization for parallelizing within a single span if the
// TBroadcastHelper instance was setup to enable that.
template <typename TBroadcastHelper>
void BroadcastLooper(TBroadcastHelper& helper, const ProcessBroadcastSpanFuncs& functors) {
ORT_ENFORCE(helper.HaveTwoTensorInputs(), "BroadcastLooper requires two tensors as input.");
if (helper.Threadpool() != nullptr && helper.SingleSpanOutput()) {
ParallelizeSingleSpan(helper, functors);
} else {
if (helper.IsInput0Scalar()) {
while (helper.NeedMoreOutput()) {
functors.input0scalar(helper);
helper.Next();
}
} else if (helper.IsInput1Scalar()) {
while (helper.NeedMoreOutput()) {
functors.input1scalar(helper);
helper.Next();
}
} else {
while (helper.NeedMoreOutput()) {
functors.general(helper);
helper.Next();
}
}
}
}
struct TensorAllocator {
TensorAllocator(OpKernelContext& context) {
auto status = context.GetTempSpaceAllocator(&allocator_);
ORT_ENFORCE(status.IsOK(), status.ErrorMessage());
}
std::unique_ptr<Tensor> Allocate(const TensorShape& shape) {
template <typename T>
std::unique_ptr<Tensor> Allocate(const TensorShape& shape) const {
return onnxruntime::make_unique<Tensor>(DataTypeImpl::GetType<T>(),
shape,
allocator_);
@ -721,161 +963,4 @@ struct TensorAllocator {
private:
AllocatorPtr allocator_;
};
// Broadcast loop for when using eigen, functions are in this form:
// Input0Scalar: [](EigenVectorMap<TOutput> output, TInput0 input0, ConstEigenVectorMap<TInput1> input1)
// Input1Scalar: [](EigenVectorMap<TOutput> output, ConstEigenVectorMap<TInput0> input0, TInput1 input1)
// General : [](EigenVectorMap<TOutput> output, ConstEigenVectorMap<TInput0> input0,
// ConstEigenVectorMap<TInput1> input1)
// Scalar parameters can also be of type const TX&.
template <typename TBroadcaster, typename Output, typename Input0Scalar, typename Input1Scalar, typename General>
void BroadcastLoop(TBroadcaster& bc, Output& output, Input0Scalar input0scalar, Input1Scalar input1scalar, General general) {
if (bc.IsInput0Scalar()) {
while (output)
input0scalar(output.NextEigenOutput(), bc.NextScalar0(), bc.NextEigen1());
} else if (bc.IsInput1Scalar()) {
while (output)
input1scalar(output.NextEigenOutput(), bc.NextEigen0(), bc.NextScalar1());
} else {
while (output)
general(output.NextEigenOutput(), bc.NextEigen0(), bc.NextEigen1());
}
}
// Broadcast loop for when using gsl::span<T>, functions are in this form:
// Input0Scalar: [](gsl::span<TOutput> output, TInput0 input0, gsl::span<const TInput1> input1)
// Input1Scalar: [](gsl::span<TOutput> output, gsl::span<const TInput0> input0, TInput1 input1)
// General : [](gsl::span<TOutput> output, gsl::span<const TInput0> input0, gsl::span<const TInput1> input1)
// Scalar parameters can also be of type const TX&.
template <typename TBroadcaster, typename Output, typename Input0Scalar, typename Input1Scalar, typename General>
void BroadcastLoopSpan(TBroadcaster& bc, Output& output, Input0Scalar input0scalar, Input1Scalar input1scalar, General general) {
if (bc.IsInput0Scalar()) {
while (output)
input0scalar(output.NextSpanOutput(), bc.NextScalar0(), bc.NextSpan1());
} else if (bc.IsInput1Scalar()) {
while (output)
input1scalar(output.NextSpanOutput(), bc.NextSpan0(), bc.NextScalar1());
} else {
while (output)
general(output.NextSpanOutput(), bc.NextSpan0(), bc.NextSpan1());
}
}
template <typename TInput, typename TOutput, typename Input0Scalar, typename Input1Scalar, typename General>
void BroadcastOneSpan(concurrency::ThreadPool* tp, double unit_cost, TOutput* output_ptr, int64_t output_size,
const TInput* input0_ptr, int64_t input0_size, const TInput* input1_ptr, int64_t input1_size,
Input0Scalar input0scalar, Input1Scalar input1scalar, General general) {
if (input0_size == 1) {
ORT_ENFORCE(input1_size == output_size);
concurrency::ThreadPool::TryParallelFor(tp, output_size,
{static_cast<float>(sizeof(TInput)), static_cast<float>(sizeof(TOutput)), unit_cost},
[=](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
EigenVectorMap<TOutput> output_map(output_ptr + first, count);
ConstEigenVectorMap<TInput> input1_map(input1_ptr + first, count);
input0scalar(output_map, *input0_ptr, input1_map);
});
} else if (input1_size == 1) {
ORT_ENFORCE(input0_size == output_size);
concurrency::ThreadPool::TryParallelFor(tp, output_size,
{static_cast<float>(sizeof(TInput)), static_cast<float>(sizeof(TOutput)), unit_cost},
[=](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
EigenVectorMap<TOutput> output_map(output_ptr + first, count);
ConstEigenVectorMap<TInput> input0_map(input0_ptr + first, count);
input1scalar(output_map, input0_map, *input1_ptr);
});
} else {
concurrency::ThreadPool::TryParallelFor(tp, output_size,
{static_cast<float>(sizeof(TInput)), static_cast<float>(sizeof(TOutput)), unit_cost},
[=](std::ptrdiff_t first, std::ptrdiff_t last) {
size_t count = static_cast<size_t>(last - first);
EigenVectorMap<TOutput> output_map(output_ptr + first, count);
ConstEigenVectorMap<TInput> input0_map(input0_ptr + first, count);
ConstEigenVectorMap<TInput> input1_map(input1_ptr + first, count);
general(output_map, input0_map, input1_map);
});
}
}
template <typename TInput, typename TOutput, typename Input0Scalar, typename Input1Scalar, typename General>
Status BroadcastTwo(OpKernelContext& context, Input0Scalar input0scalar, Input1Scalar input1scalar, General general, double unit_cost=-1.0f) {
if (unit_cost == -1.0f) { // no paralellization
TBroadcaster<TInput, TInput> bc(*context.Input<Tensor>(0), *context.Input<Tensor>(1));
TBroadcastOutput<TOutput> output(bc.GetSpanSize(), *context.Output(0, bc.GetOutputShape()));
BroadcastLoop(bc, output, input0scalar, input1scalar, general);
} else {
const Tensor* input0_tensor = context.Input<Tensor>(0);
const Tensor* input1_tensor = context.Input<Tensor>(1);
TBroadcaster<TInput, TInput> bc(*input0_tensor, *input1_tensor);
Tensor& output_tensor = *context.Output(0, bc.GetOutputShape());
auto span_size = bc.GetSpanSize();
int64_t output_size = output_tensor.Shape().Size();
if (output_size != 0) {
concurrency::ThreadPool* tp = context.GetOperatorThreadPool();
if (span_size != 0) {
if (output_size == static_cast<int64_t>(span_size)) { // Only one big span for all data, parallel inside it
ORT_ENFORCE((output_size % span_size) == 0);
BroadcastOneSpan(tp, unit_cost, output_tensor.MutableData<TOutput>(), output_size,
input0_tensor->Data<TInput>(), input0_tensor->Shape().Size(),
input1_tensor->Data<TInput>(), input1_tensor->Shape().Size(),
input0scalar, input1scalar, general);
} else {
concurrency::ThreadPool::TryParallelFor(
tp, output_size / span_size,
{static_cast<float>(sizeof(TInput)) * span_size, static_cast<float>(sizeof(TOutput)) * span_size, unit_cost * span_size},
[=, &bc, &output_tensor](std::ptrdiff_t first_span, std::ptrdiff_t last_span) {
TBroadcaster<TInput, TInput> span_bc(bc);
TBroadcastOutput<TOutput> span_output(span_size, output_tensor, first_span * span_size, last_span * span_size);
span_bc.AdvanceBy(first_span * span_size);
BroadcastLoop(span_bc, span_output, input0scalar, input1scalar, general);
});
}
}
}
}
return Status::OK();
}
template <typename TInput, typename TOutput, typename Input0Scalar, typename Input1Scalar, typename General>
Status BroadcastVariadic(const Node& node, OpKernelContext& context, Input0Scalar input0scalar, Input1Scalar input1scalar, General general) {
auto input_count = node.InputArgCount().front();
ORT_ENFORCE(input_count >= 1, "Must have 1 or more inputs");
// One item is trivial, just copy across and exit
if (input_count == 1) {
EigenMap<TOutput>(*context.Output(0, context.Input<Tensor>(0)->Shape())) = EigenMap<TInput>(*context.Input<Tensor>(0));
return Status::OK();
}
std::unique_ptr<Tensor> tempInput;
std::unique_ptr<Tensor> tempOutput;
TensorAllocator<TOutput> tensorAllocator(context);
// For more than 2 tensors, we sum the first two into a temporary tensor, then sum the next with the temporary tensor
for (int i = 0; i < input_count - 1; i++) {
auto& tensor0 = tempInput ? *tempInput : *context.Input<Tensor>(0);
auto& tensor1 = *context.Input<Tensor>(i + 1);
TBroadcaster<TInput, TInput> bc(tensor0, tensor1);
// Create a temporary output for all but the last iteration, which goes to the real output
Tensor* p_output{};
if (i == input_count - 2)
p_output = context.Output(0, bc.GetOutputShape());
else {
tempOutput = tensorAllocator.Allocate(bc.GetOutputShape());
p_output = tempOutput.get();
}
TBroadcastOutput<TOutput> output(bc.GetSpanSize(), *p_output);
BroadcastLoop(bc, output, input0scalar, input1scalar, general);
tempInput = std::move(tempOutput);
}
return Status::OK();
}
} // namespace onnxruntime

View file

@ -42,136 +42,189 @@ WHERE_TYPED_KERNEL_WITH_TYPE_NAME(std::string, string)
namespace {
template<typename T, typename R>
template <typename T, typename R>
using EnableIfEigenScalar = typename std::enable_if<std::is_arithmetic<T>::value, R>::type;
template <typename T, typename R>
using EnableIfEigenNotScalar = typename std::enable_if<!std::is_arithmetic<T>::value, R>::type;
template <typename T>
EnableIfEigenScalar<T, void>
SelectBroadcastLoop(bool target,
TBroadcaster<bool, T>* select_broadcaster,
TBroadcastOutput<T>* select_broadcast_output) {
BroadcastLoop(
*select_broadcaster, *select_broadcast_output,
[target](EigenVectorMap<T> output, bool condition, ConstEigenVectorMap<T> value) {
EnableIfEigenScalar<T, ProcessBroadcastSpanFuncs> SelectBroadcastFuncs() {
return ProcessBroadcastSpanFuncs{
[](BroadcastHelper& per_iter_bh) {
bool target = per_iter_bh.GetUserData();
bool condition = per_iter_bh.ScalarInput0<bool>();
auto value = per_iter_bh.EigenInput1<T>();
auto output = per_iter_bh.OutputEigen<T>();
if (condition == target) {
output = value;
} else {
output = EigenVectorMap<T>::PlainObject::Constant(value.size(), T{});
}
},
[target](EigenVectorMap<T> output, ConstEigenVectorMap<bool> condition, const T& value) {
[](BroadcastHelper& per_iter_bh) {
bool target = per_iter_bh.GetUserData();
auto condition = per_iter_bh.EigenInput0<bool>();
const T& value = per_iter_bh.ScalarInput1<T>();
auto output = per_iter_bh.OutputEigen<T>();
output = (condition.array() == target)
.select(value, EigenVectorMap<T>::PlainObject::Constant(condition.size(), T{}));
},
[target](EigenVectorMap<T> output, ConstEigenVectorMap<bool> condition, ConstEigenVectorMap<T> value) {
[](BroadcastHelper& per_iter_bh) {
bool target = per_iter_bh.GetUserData();
auto condition = per_iter_bh.EigenInput0<bool>();
auto value = per_iter_bh.EigenInput1<T>();
auto output = per_iter_bh.OutputEigen<T>();
output = (condition.array() == target)
.select(value, EigenVectorMap<T>::PlainObject::Constant(condition.size(), T{}));
});
}};
}
template <typename T>
EnableIfEigenNotScalar<T, void>
SelectBroadcastLoop(bool target, TBroadcaster<bool, T>* select_broadcaster,
TBroadcastOutput<T>* select_broadcast_output) {
BroadcastLoopSpan(
*select_broadcaster, *select_broadcast_output,
[target](gsl::span<T> output, bool condition, gsl::span<const T> value) {
EnableIfEigenNotScalar<T, ProcessBroadcastSpanFuncs> SelectBroadcastFuncs() {
return ProcessBroadcastSpanFuncs{
[](BroadcastHelper& per_iter_bh) {
bool target = per_iter_bh.GetUserData();
bool condition = per_iter_bh.ScalarInput0<bool>();
auto value = per_iter_bh.SpanInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
if (condition == target) {
std::copy(value.cbegin(), value.cend(), output.begin());
} else {
std::fill(output.begin(), output.end(), T{});
}
},
[target](gsl::span<T> output, gsl::span<const bool> condition, const T& value) {
[](BroadcastHelper& per_iter_bh) {
bool target = per_iter_bh.GetUserData();
auto condition = per_iter_bh.SpanInput0<bool>();
const T& value = per_iter_bh.ScalarInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
std::transform(condition.cbegin(), condition.cend(), output.begin(),
[target, &value](bool condition_element) {
return condition_element == target ? value : T{};
});
},
[target](gsl::span<T> output, gsl::span<const bool> condition, gsl::span<const T> value) {
[](BroadcastHelper& per_iter_bh) {
bool target = per_iter_bh.GetUserData();
auto condition = per_iter_bh.SpanInput0<bool>();
auto value = per_iter_bh.SpanInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
std::transform(condition.cbegin(), condition.cend(), value.cbegin(), output.begin(),
[target](bool condition_element, const T& value_element) {
return condition_element == target ? value_element : T{};
});
});
}};
}
template <typename T>
std::unique_ptr<Tensor> Select(bool target, const Tensor& condition_tensor, const Tensor& value_tensor,
TensorAllocator<T>& tensor_allocator) {
TBroadcaster<bool, T> select_broadcaster{condition_tensor, value_tensor};
std::unique_ptr<Tensor> select_tensor{
tensor_allocator.Allocate(select_broadcaster.GetOutputShape())};
TBroadcastOutput<T> select_broadcast_output{
select_broadcaster.GetSpanSize(), *select_tensor};
void MergeScalarAndVector(EigenVectorMap<T> output, const T& scalar_value, ConstEigenVectorMap<T> vector_value) {
if (scalar_value != T{}) {
output = EigenVectorMap<T>::PlainObject::Constant(vector_value.size(), scalar_value);
} else {
output = vector_value;
}
};
SelectBroadcastLoop(target, &select_broadcaster, &select_broadcast_output);
return select_tensor;
template <typename T>
EnableIfEigenScalar<T, ProcessBroadcastSpanFuncs> MergeBroadcastFuncs() {
return ProcessBroadcastSpanFuncs{
[](BroadcastHelper& per_iter_bh) {
MergeScalarAndVector(per_iter_bh.OutputEigen<T>(),
per_iter_bh.ScalarInput0<T>(), // X selection
per_iter_bh.EigenInput1<T>()); // Y selection
},
[](BroadcastHelper& per_iter_bh) {
MergeScalarAndVector(per_iter_bh.OutputEigen<T>(),
per_iter_bh.ScalarInput1<T>(), // Y selection
per_iter_bh.EigenInput0<T>()); // X selection
},
[](BroadcastHelper& per_iter_bh) {
auto X_selection = per_iter_bh.EigenInput0<T>();
auto Y_selection = per_iter_bh.EigenInput1<T>();
per_iter_bh.OutputEigen<T>() = X_selection.binaryExpr(Y_selection,
[](T x, T y) -> T {
return x != T{} ? x : y;
});
}};
}
template <typename T>
EnableIfEigenScalar<T, void>
MergeBroadcastLoop(TBroadcaster<T, T>* merge_broadcaster, TBroadcastOutput<T>* merge_broadcast_output) {
const auto merge_scalar_and_vector = [](EigenVectorMap<T> output,
const T& scalar_value, ConstEigenVectorMap<T> vector_value) {
if (scalar_value != T{}) {
output = EigenVectorMap<T>::PlainObject::Constant(vector_value.size(), scalar_value);
} else {
output = vector_value;
}
};
BroadcastLoop(
*merge_broadcaster, *merge_broadcast_output,
[merge_scalar_and_vector](EigenVectorMap<T> output, const T& X_selection, ConstEigenVectorMap<T> Y_selection) {
merge_scalar_and_vector(output, X_selection, Y_selection);
},
[merge_scalar_and_vector](EigenVectorMap<T> output, ConstEigenVectorMap<T> X_selection, const T& Y_selection) {
merge_scalar_and_vector(output, Y_selection, X_selection);
},
[](EigenVectorMap<T> output, ConstEigenVectorMap<T> X_selection, ConstEigenVectorMap<T> Y_selection) {
output = X_selection.binaryExpr(Y_selection, [](T x, T y) -> T {
return x != T{} ? x : y;
});
});
}
void MergeScalarAndVector(gsl::span<T> output, const T& scalar_value, gsl::span<const T> vector_value) {
if (!scalar_value.empty()) {
std::fill(output.begin(), output.end(), scalar_value);
} else {
std::copy(vector_value.cbegin(), vector_value.cend(), output.begin());
}
};
template <typename T>
EnableIfEigenNotScalar<T, void>
MergeBroadcastLoop(TBroadcaster<T, T>* merge_broadcaster, TBroadcastOutput<T>* merge_broadcast_output) {
const auto merge_scalar_and_vector = [](gsl::span<T> output, const T& scalar_value, gsl::span<const T> vector_value) {
if (!scalar_value.empty()) {
std::fill(output.begin(), output.end(), scalar_value);
} else {
std::copy(vector_value.cbegin(), vector_value.cend(), output.begin());
}
};
BroadcastLoopSpan(
*merge_broadcaster, *merge_broadcast_output,
[merge_scalar_and_vector](gsl::span<T> output, const T& X_selection, gsl::span<const T> Y_selection) {
merge_scalar_and_vector(output, X_selection, Y_selection);
EnableIfEigenNotScalar<T, ProcessBroadcastSpanFuncs> MergeBroadcastFuncs() {
return ProcessBroadcastSpanFuncs{
[](BroadcastHelper& per_iter_bh) {
MergeScalarAndVector(per_iter_bh.OutputSpan<T>(),
per_iter_bh.ScalarInput0<T>(), // X selection
per_iter_bh.SpanInput1<T>()); // Y selection
},
[merge_scalar_and_vector](gsl::span<T> output, gsl::span<const T> X_selection, const T& Y_selection) {
merge_scalar_and_vector(output, Y_selection, X_selection);
[](BroadcastHelper& per_iter_bh) {
MergeScalarAndVector(per_iter_bh.OutputSpan<T>(),
per_iter_bh.ScalarInput1<T>(), // Y selection
per_iter_bh.SpanInput0<T>()); // X selection
},
[](gsl::span<T> output, gsl::span<const T> X_selection, gsl::span<const T> Y_selection) {
[](BroadcastHelper& per_iter_bh) {
auto X_selection = per_iter_bh.SpanInput0<T>();
auto Y_selection = per_iter_bh.SpanInput1<T>();
auto output = per_iter_bh.OutputSpan<T>();
std::transform(X_selection.cbegin(), X_selection.cend(), Y_selection.cbegin(), output.begin(),
[](const T& x, const T& y) { return !x.empty() ? x : y; });
});
}};
}
// function pointer to create typed tensor from type agnostic code whilst avoiding the overhead of std::function
using AllocTensorFunc = std::unique_ptr<Tensor> (*)(const TensorAllocator& allocator, const TensorShape& shape);
static std::unique_ptr<Tensor> UntypedSelect(OpKernelContext& context, bool target,
const TensorAllocator& allocator, AllocTensorFunc allocate_tensor,
const ProcessBroadcastSpanFuncs& functors) {
const auto& condition = *context.Input<Tensor>(0);
// select the X input (input 1) for 'true', and Y input (input 2) for 'false'
const auto& values = *context.Input<Tensor>(target ? 1 : 2);
InputBroadcaster input_broadcaster(condition, values);
std::unique_ptr<Tensor> selection_tensor = allocate_tensor(allocator, input_broadcaster.GetOutputShape());
OutputBroadcaster output_broadcaster(input_broadcaster.GetSpanSize(), *selection_tensor);
// store value of 'target' directly in void* for user_data so it's accessible in the state-less functors
BroadcastHelper broadcast_helper(input_broadcaster, output_broadcaster, reinterpret_cast<void*>(target));
BroadcastLooper(broadcast_helper, functors);
return selection_tensor;
}
static void UntypedMerge(OpKernelContext& context,
const Tensor& X_selection_tensor, const Tensor& Y_selection_tensor,
const ProcessBroadcastSpanFuncs& functors) {
InputBroadcaster merge_broadcaster{X_selection_tensor, Y_selection_tensor};
Tensor& output = *context.Output(0, merge_broadcaster.GetOutputShape());
OutputBroadcaster output_broadcaster{merge_broadcaster.GetSpanSize(), output};
BroadcastHelper broadcast_helper(merge_broadcaster, output_broadcaster);
BroadcastLooper(broadcast_helper, functors);
}
} // namespace
template <typename T>
Status Where<T>::Compute(OpKernelContext* context) const {
const auto* const condition = context->Input<Tensor>(0);
const auto* const X = context->Input<Tensor>(1);
const auto* const Y = context->Input<Tensor>(2);
ORT_ENFORCE(condition && X && Y, "condition, X, and Y inputs are required!");
// we use a func pointer to save the overhead of std::function, so we can't capture tensor_allocator here
const auto typed_tensor_allocation = [](const TensorAllocator& allocator,
const TensorShape& shape) {
return allocator.Allocate<T>(shape);
};
TensorAllocator tensor_allocator{*context};
ProcessBroadcastSpanFuncs funcs = SelectBroadcastFuncs<T>();
// The current implementation is limited to broadcasting over two tensors at once.
// So, we first broadcast over condition and X to select the values from X:
@ -180,17 +233,10 @@ Status Where<T>::Compute(OpKernelContext* context) const {
// Y_selection = !condition ? Y : default value
// Finally, we broadcast over and merge X_selection and Y_selection:
// output = (X_selection != default value) ? X_selection : Y_selection
TensorAllocator<T> tensor_allocator{*context};
auto X_selection_tensor = Select<T>(true, *condition, *X, tensor_allocator);
auto Y_selection_tensor = Select<T>(false, *condition, *Y, tensor_allocator);
auto X_selection_tensor = UntypedSelect(*context, true, tensor_allocator, typed_tensor_allocation, funcs);
auto Y_selection_tensor = UntypedSelect(*context, false, tensor_allocator, typed_tensor_allocation, funcs);
TBroadcaster<T, T> merge_broadcaster{*X_selection_tensor, *Y_selection_tensor};
Tensor* const output = context->Output(0, merge_broadcaster.GetOutputShape());
ORT_ENFORCE(output, "failed to get first output!");
TBroadcastOutput<T> merge_broadcast_output{
merge_broadcaster.GetSpanSize(), *output};
MergeBroadcastLoop(&merge_broadcaster, &merge_broadcast_output);
UntypedMerge(*context, *X_selection_tensor, *Y_selection_tensor, MergeBroadcastFuncs<T>());
return Status::OK();
}

View file

@ -37,12 +37,22 @@ Status InPlaceAccumulator<T>::Compute(OpKernelContext* context) const {
}
//Copy from Add CPU kernel
return BroadcastTwo<T, T>(
*context,
[](EigenVectorMap<T> output, T input0, ConstEigenVectorMap<T> input1) { output = input0 + input1.array(); },
[](EigenVectorMap<T> output, ConstEigenVectorMap<T> input0, T input1) { output = input0.array() + input1; },
[](EigenVectorMap<T> output, ConstEigenVectorMap<T> input0, ConstEigenVectorMap<T> input1) { output = input0 + input1; });
ProcessBroadcastSpanFuncs funcs{
[](BroadcastHelper& per_iter_bh) {
per_iter_bh.OutputEigen<T>() = per_iter_bh.ScalarInput0<T>() + per_iter_bh.EigenInput1<T>().array();
},
[](BroadcastHelper& per_iter_bh) {
per_iter_bh.OutputEigen<T>() = per_iter_bh.EigenInput0<T>().array() + per_iter_bh.ScalarInput1<T>();
},
[](BroadcastHelper& per_iter_bh) {
per_iter_bh.OutputEigen<T>() = per_iter_bh.EigenInput0<T>() + per_iter_bh.EigenInput1<T>();
}};
UntypedBroadcastTwo(*context, funcs);
return Status::OK();
}
template <typename T>
Status ZeroGradient<T>::Compute(OpKernelContext* context) const {
const Tensor& old_gradient = *context->Input<Tensor>(0);