mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
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:
parent
43faf9e388
commit
c52561d044
5 changed files with 1161 additions and 706 deletions
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
Loading…
Reference in a new issue