mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Updated SoftmaxGrad and LogSoftmaxGrad to support version 13. (#9733)
* Updated SoftmaxGrad_13/LogSoftmaxGrad_13 to support version 13.
This commit is contained in:
parent
dc1724b0e2
commit
3af14fc554
10 changed files with 269 additions and 85 deletions
|
|
@ -99,6 +99,10 @@
|
|||
"LogSoftmaxGrad com.microsoft CPUExecutionProvider",
|
||||
2657523710083167200
|
||||
],
|
||||
[
|
||||
"LogSoftmaxGrad_13 com.microsoft CPUExecutionProvider",
|
||||
1917456134240183096
|
||||
],
|
||||
[
|
||||
"MaxPoolGrad ai.onnx CPUExecutionProvider",
|
||||
17526822836083413768
|
||||
|
|
@ -239,6 +243,10 @@
|
|||
"SoftmaxGrad com.microsoft CPUExecutionProvider",
|
||||
4483165757863027152
|
||||
],
|
||||
[
|
||||
"SoftmaxGrad_13 com.microsoft CPUExecutionProvider",
|
||||
8375491041422269560
|
||||
],
|
||||
[
|
||||
"SparseSoftmaxCrossEntropy ai.onnx CPUExecutionProvider",
|
||||
10638058507241762520
|
||||
|
|
|
|||
|
|
@ -707,7 +707,7 @@ IMPLEMENT_GRADIENT_BUILDER(GetSigmoidGradient) {
|
|||
|
||||
IMPLEMENT_GRADIENT_BUILDER(GetSoftmaxGradient) {
|
||||
return std::vector<NodeDef>{
|
||||
NodeDef(OpDef{"SoftmaxGrad", kMSDomain, 1},
|
||||
NodeDef(OpDef{SrcNodeOpsetVersion() < 13 ? "SoftmaxGrad" : "SoftmaxGrad_13", kMSDomain, 1},
|
||||
{GO(0), O(0)},
|
||||
{GI(0)},
|
||||
SrcNodeAttributes())};
|
||||
|
|
@ -715,7 +715,7 @@ IMPLEMENT_GRADIENT_BUILDER(GetSoftmaxGradient) {
|
|||
|
||||
IMPLEMENT_GRADIENT_BUILDER(GetLogSoftmaxGradient) {
|
||||
return std::vector<NodeDef>{
|
||||
NodeDef(OpDef{"LogSoftmaxGrad", kMSDomain, 1},
|
||||
NodeDef(OpDef{SrcNodeOpsetVersion() < 13 ? "LogSoftmaxGrad" : "LogSoftmaxGrad_13", kMSDomain, 1},
|
||||
{GO(0), O(0)},
|
||||
{GI(0)},
|
||||
SrcNodeAttributes())};
|
||||
|
|
|
|||
|
|
@ -650,6 +650,24 @@ void RegisterTrainingOpSchemas() {
|
|||
return true;
|
||||
});
|
||||
|
||||
ONNX_CONTRIB_OPERATOR_SCHEMA(SoftmaxGrad_13)
|
||||
.SetDomain(kMSDomain)
|
||||
.SinceVersion(1)
|
||||
.Input(0, "dY", "Gradient of output Y", "T")
|
||||
.Input(1, "Y", "Input tensor", "T")
|
||||
.Output(0, "dX", "Gradient of input X", "T")
|
||||
.Attr(
|
||||
"axis",
|
||||
"Describes the dimension Softmax will be performed on."
|
||||
"Defaults to -1. Negative value means counting dimensions from the back.",
|
||||
AttributeProto::INT,
|
||||
static_cast<int64_t>(-1))
|
||||
.TypeConstraint(
|
||||
"T",
|
||||
{"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"},
|
||||
"Constrain input and output types to float tensors.")
|
||||
.TypeAndShapeInferenceFunction(propagateShapeAndTypeFromFirstInput);
|
||||
|
||||
ONNX_CONTRIB_OPERATOR_SCHEMA(LogSoftmaxGrad)
|
||||
.SetDomain(kMSDomain)
|
||||
.SinceVersion(1)
|
||||
|
|
@ -669,6 +687,24 @@ void RegisterTrainingOpSchemas() {
|
|||
"Constrain input and output types to float tensors.")
|
||||
.TypeAndShapeInferenceFunction(propagateShapeAndTypeFromFirstInput);
|
||||
|
||||
ONNX_CONTRIB_OPERATOR_SCHEMA(LogSoftmaxGrad_13)
|
||||
.SetDomain(kMSDomain)
|
||||
.SinceVersion(1)
|
||||
.Input(0, "dY", "Gradient of output Y", "T")
|
||||
.Input(1, "X", "Input tensor", "T")
|
||||
.Output(0, "dX", "Gradient of input X", "T")
|
||||
.Attr(
|
||||
"axis",
|
||||
"Describes the dimension LogSoftmax will be performed on."
|
||||
"Defaults to -1. Negative value means counting dimensions from the back.",
|
||||
AttributeProto::INT,
|
||||
static_cast<int64_t>(-1))
|
||||
.TypeConstraint(
|
||||
"T",
|
||||
{"tensor(float16)", "tensor(float)", "tensor(double)"},
|
||||
"Constrain input and output types to float tensors.")
|
||||
.TypeAndShapeInferenceFunction(propagateShapeAndTypeFromFirstInput);
|
||||
|
||||
ONNX_CONTRIB_OPERATOR_SCHEMA(AveragePoolGrad)
|
||||
.SinceVersion(9)
|
||||
.Input(0, "dY", "Gradient of output Y", "T")
|
||||
|
|
|
|||
|
|
@ -1574,13 +1574,13 @@ TEST(GradientCheckerTest, SigmoidGrad) {
|
|||
UnaryOpGradientTest("Sigmoid");
|
||||
}
|
||||
|
||||
void GradientCheckerSoftmaxGradHelper(bool is_log_softmax) {
|
||||
void GradientCheckerSoftmaxGradHelper(bool is_log_softmax, int version = 11) {
|
||||
TensorShape shape({3, 4, 5});
|
||||
float max_error;
|
||||
GradientChecker<float, float, float> gradient_checker;
|
||||
|
||||
const std::string op = is_log_softmax ? "LogSoftmax" : "Softmax";
|
||||
OpDef op_def{op};
|
||||
OpDef op_def{op, kOnnxDomain, version};
|
||||
|
||||
// default_axis
|
||||
{
|
||||
|
|
@ -1594,6 +1594,12 @@ void GradientCheckerSoftmaxGradHelper(bool is_log_softmax) {
|
|||
EXPECT_IS_TINY(max_error);
|
||||
}
|
||||
|
||||
// axis=1
|
||||
{
|
||||
ASSERT_STATUS_OK(gradient_checker.ComputeGradientError(op_def, {shape}, {shape}, &max_error, {MakeAttribute("axis", int64_t(1))}));
|
||||
EXPECT_IS_TINY(max_error);
|
||||
}
|
||||
|
||||
// axis=2
|
||||
{
|
||||
ASSERT_STATUS_OK(gradient_checker.ComputeGradientError(op_def, {shape}, {shape}, &max_error, {MakeAttribute("axis", int64_t(2))}));
|
||||
|
|
@ -1603,10 +1609,12 @@ void GradientCheckerSoftmaxGradHelper(bool is_log_softmax) {
|
|||
|
||||
TEST(GradientCheckerTest, SoftMaxGrad) {
|
||||
GradientCheckerSoftmaxGradHelper(false);
|
||||
GradientCheckerSoftmaxGradHelper(false, 13);
|
||||
}
|
||||
|
||||
TEST(GradientCheckerTest, LogSoftMaxGrad) {
|
||||
GradientCheckerSoftmaxGradHelper(true);
|
||||
GradientCheckerSoftmaxGradHelper(true, 13);
|
||||
}
|
||||
|
||||
void TestSoftmaxCrossEntropyGrad(const TensorShape& input_shape, const std::string& reduction) {
|
||||
|
|
|
|||
|
|
@ -40,6 +40,8 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, ConvG
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, ReluGrad);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxGrad);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, LogSoftmaxGrad);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, LogSoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, AveragePoolGrad);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MaxPoolGrad);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherGrad);
|
||||
|
|
@ -149,6 +151,8 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, ReluGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, LogSoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxGrad_13)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, LogSoftmaxGrad_13)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, AveragePoolGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MaxPoolGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherGrad)>,
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@
|
|||
#include "core/util/math.h"
|
||||
#include "core/providers/cpu/math/element_wise_ops.h"
|
||||
#include "core/providers/cpu/math/matmul_helper.h"
|
||||
#include "core/providers/cpu/tensor/transpose.h"
|
||||
#include "gsl/gsl"
|
||||
|
||||
namespace onnxruntime {
|
||||
|
|
@ -58,6 +59,14 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
SoftmaxGrad<float>);
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
SoftmaxGrad_13,
|
||||
kMSDomain,
|
||||
1,
|
||||
kCpuExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
SoftmaxGrad<float>);
|
||||
|
||||
template <typename T>
|
||||
Status SoftmaxGrad<T>::Compute(OpKernelContext* context) const {
|
||||
auto& dY = *context->Input<Tensor>(0);
|
||||
|
|
@ -65,7 +74,8 @@ Status SoftmaxGrad<T>::Compute(OpKernelContext* context) const {
|
|||
const TensorShape input_shape{Y.Shape()};
|
||||
auto& dX = *context->Output(0, Y.Shape());
|
||||
|
||||
auto axis = HandleNegativeAxis(axis_, Y.Shape().NumDimensions());
|
||||
size_t rank = input_shape.NumDimensions();
|
||||
const size_t axis = static_cast<size_t>(HandleNegativeAxis(axis_, rank));
|
||||
|
||||
size_t N = input_shape.SizeToDimension(axis);
|
||||
size_t D = input_shape.SizeFromDimension(axis);
|
||||
|
|
@ -74,30 +84,88 @@ Status SoftmaxGrad<T>::Compute(OpKernelContext* context) const {
|
|||
return Status::OK();
|
||||
}
|
||||
|
||||
std::vector<float> scale_(N);
|
||||
std::vector<float> sum_multiplier_(D, 1.f); // initialize all multiplier values to 1.0
|
||||
bool is_transpose_required = opset_ >= 13 && axis != (rank - 1);
|
||||
|
||||
std::unique_ptr<Tensor> transposed_dY;
|
||||
std::unique_ptr<Tensor> transposed_Y;
|
||||
std::vector<int64_t> transposed_input_dims;
|
||||
std::unique_ptr<Tensor> intermediate_output; // output that the softmax implementation will write into while using transposed input
|
||||
std::vector<size_t> permutation(rank);
|
||||
|
||||
if (is_transpose_required) {
|
||||
AllocatorPtr alloc;
|
||||
auto status = context->GetTempSpaceAllocator(&alloc);
|
||||
if (!status.IsOK())
|
||||
return status;
|
||||
|
||||
std::iota(std::begin(permutation), std::end(permutation), 0);
|
||||
|
||||
// swap the innermost dim with the dim corresponding to axis
|
||||
permutation[axis] = rank - 1;
|
||||
permutation[rank - 1] = axis;
|
||||
|
||||
transposed_input_dims.reserve(rank);
|
||||
for (auto e : permutation) {
|
||||
transposed_input_dims.push_back(input_shape[e]);
|
||||
}
|
||||
N = TensorShape(transposed_input_dims).SizeToDimension(rank - 1);
|
||||
D = TensorShape(transposed_input_dims).SizeFromDimension(rank - 1);
|
||||
|
||||
// Allocate a temporary tensor to hold transposed input
|
||||
auto temp_input0 = Tensor::Create(Y.DataType(), TensorShape(transposed_input_dims), alloc);
|
||||
|
||||
// Perform the transpose
|
||||
ORT_RETURN_IF_ERROR(Transpose::DoTranspose(permutation, Y, *temp_input0));
|
||||
transposed_Y = std::move(temp_input0);
|
||||
|
||||
auto temp_input1 = Tensor::Create(Y.DataType(), TensorShape(transposed_input_dims), alloc);
|
||||
ORT_RETURN_IF_ERROR(Transpose::DoTranspose(permutation, dY, *temp_input1));
|
||||
transposed_dY = std::move(temp_input1);
|
||||
|
||||
// Allocate memory for the intermediate output
|
||||
intermediate_output = Tensor::Create(dX.DataType(), TensorShape(transposed_input_dims), alloc);
|
||||
}
|
||||
|
||||
const int n = gsl::narrow_cast<int>(N);
|
||||
const int d = gsl::narrow_cast<int>(D);
|
||||
const int nd = gsl::narrow_cast<int>(N * D);
|
||||
|
||||
float* scaledata = scale_.data();
|
||||
const float* Ydata = Y.template Data<float>();
|
||||
const float* dYdata = dY.template Data<float>();
|
||||
float* dXdata = dX.template MutableData<float>();
|
||||
const float* Ydata = is_transpose_required ? transposed_Y->template Data<T>() : Y.template Data<float>();
|
||||
const float* dYdata = is_transpose_required ? transposed_dY->template Data<T>() : dY.template Data<float>();
|
||||
float* dXdata = is_transpose_required ? intermediate_output->template MutableData<T>() : dX.template MutableData<float>();
|
||||
|
||||
gsl::copy(gsl::make_span(dYdata, nd), gsl::make_span(dXdata, nd));
|
||||
if (is_logsoftmaxgrad_) {
|
||||
std::vector<float> eY(nd);
|
||||
float* eYdata = eY.data();
|
||||
|
||||
for (size_t i = 0; i < N; ++i) {
|
||||
math::Dot<float, CPUMathUtil>(d, Ydata + i * d, dYdata + i * d,
|
||||
scaledata + i, nullptr);
|
||||
// dX_ai = d(log Y_ai) - [sum_j d(log Y_aj)] exp(log Y_ai)
|
||||
gsl::copy(gsl::make_span(dYdata, nd), gsl::make_span(dXdata, nd));
|
||||
math::Exp<float, CPUMathUtil>(nd, Ydata, eYdata, nullptr);
|
||||
for (size_t i = 0; i < N; ++i) {
|
||||
float sdY;
|
||||
math::Sum<float, CPUMathUtil>(d, dYdata + i * d, &sdY, nullptr, nullptr);
|
||||
math::Axpy<float, CPUMathUtil>(d, -sdY, eYdata + i * d, dXdata + i * d, nullptr);
|
||||
}
|
||||
} else {
|
||||
std::vector<float> scale_(N);
|
||||
std::vector<float> sum_multiplier_(D, 1.f); // initialize all multiplier values to 1.0
|
||||
float* scaledata = scale_.data();
|
||||
for (size_t i = 0; i < N; ++i) {
|
||||
math::Dot<float, CPUMathUtil>(d, Ydata + i * d, dYdata + i * d,
|
||||
scaledata + i, nullptr);
|
||||
}
|
||||
|
||||
concurrency::ThreadPool* tp = context->GetOperatorThreadPool();
|
||||
math::Gemm<float>(CblasNoTrans, CblasNoTrans, n, d, 1, -1,
|
||||
scaledata, sum_multiplier_.data(), 1,
|
||||
dXdata, tp);
|
||||
|
||||
math::Mul<float, CPUMathUtil>(gsl::narrow_cast<int>(Y.Shape().Size()), dXdata, Ydata, dXdata, nullptr);
|
||||
}
|
||||
if (is_transpose_required) {
|
||||
// Perform the transpose to get the axes back to the original ordering
|
||||
ORT_RETURN_IF_ERROR(Transpose::DoTranspose(permutation, *intermediate_output, dX));
|
||||
}
|
||||
|
||||
concurrency::ThreadPool* tp = context->GetOperatorThreadPool();
|
||||
math::Gemm<float>(CblasNoTrans, CblasNoTrans, n, d, 1, -1,
|
||||
scaledata, sum_multiplier_.data(), 1,
|
||||
dXdata, tp);
|
||||
|
||||
math::Mul<float, CPUMathUtil>(gsl::narrow_cast<int>(Y.Shape().Size()), dXdata, Ydata, dXdata, nullptr);
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
|
@ -108,45 +176,15 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
1,
|
||||
kCpuExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
LogSoftmaxGrad<float>);
|
||||
SoftmaxGrad<float>);
|
||||
|
||||
template <typename T>
|
||||
Status LogSoftmaxGrad<T>::Compute(OpKernelContext* context) const {
|
||||
auto& dY = *context->Input<Tensor>(0);
|
||||
auto& Y = *context->Input<Tensor>(1);
|
||||
const TensorShape input_shape{Y.Shape()};
|
||||
auto& dX = *context->Output(0, Y.Shape());
|
||||
|
||||
auto axis = HandleNegativeAxis(axis_, Y.Shape().NumDimensions());
|
||||
|
||||
size_t N = input_shape.SizeToDimension(axis);
|
||||
size_t D = input_shape.SizeFromDimension(axis);
|
||||
|
||||
if (N == 0) {
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
const int d = gsl::narrow_cast<int>(D);
|
||||
const int nd = gsl::narrow_cast<int>(N * D);
|
||||
|
||||
const float* Ydata = Y.template Data<float>();
|
||||
const float* dYdata = dY.template Data<float>();
|
||||
float* dXdata = dX.template MutableData<float>();
|
||||
|
||||
std::vector<float> eY(nd);
|
||||
float* eYdata = eY.data();
|
||||
|
||||
// dX_ai = d(log Y_ai) - [sum_j d(log Y_aj)] exp(log Y_ai)
|
||||
gsl::copy(gsl::make_span(dYdata, nd), gsl::make_span(dXdata, nd));
|
||||
math::Exp<float, CPUMathUtil>(nd, Ydata, eYdata, nullptr);
|
||||
for (size_t i = 0; i < N; ++i) {
|
||||
float sdY;
|
||||
math::Sum<float, CPUMathUtil>(d, dYdata + i * d, &sdY, nullptr, nullptr);
|
||||
math::Axpy<float, CPUMathUtil>(d, -sdY, eYdata + i * d, dXdata + i * d, nullptr);
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
LogSoftmaxGrad_13,
|
||||
kMSDomain,
|
||||
1,
|
||||
kCpuExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
SoftmaxGrad<float>);
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
SigmoidGrad,
|
||||
|
|
|
|||
|
|
@ -61,7 +61,10 @@ template <typename T>
|
|||
class SoftmaxGrad final : public OpKernel {
|
||||
public:
|
||||
explicit SoftmaxGrad(const OpKernelInfo& info) : OpKernel(info) {
|
||||
axis_ = info.GetAttrOrDefault<int64_t>("axis", 0);
|
||||
const auto& node = info.node();
|
||||
opset_ = (node.OpType() == "SoftmaxGrad_13" || node.OpType() == "LogSoftmaxGrad_13") ? 13 : 1;
|
||||
axis_ = info.GetAttrOrDefault("axis", static_cast<int64_t>(opset_ < 13 ? 1 : -1));
|
||||
is_logsoftmaxgrad_ = node.OpType() == "LogSoftmaxGrad_13" || node.OpType() == "LogSoftmaxGrad";
|
||||
}
|
||||
|
||||
Status Compute(OpKernelContext* context) const override;
|
||||
|
|
@ -69,20 +72,8 @@ class SoftmaxGrad final : public OpKernel {
|
|||
private:
|
||||
ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(SoftmaxGrad);
|
||||
int64_t axis_;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
class LogSoftmaxGrad final : public OpKernel {
|
||||
public:
|
||||
explicit LogSoftmaxGrad(const OpKernelInfo& info) : OpKernel(info) {
|
||||
axis_ = info.GetAttrOrDefault<int64_t>("axis", 0);
|
||||
}
|
||||
|
||||
Status Compute(OpKernelContext* context) const override;
|
||||
|
||||
private:
|
||||
ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(LogSoftmaxGrad);
|
||||
int64_t axis_;
|
||||
int opset_; // opset_ of the forward Softmax operator
|
||||
bool is_logsoftmaxgrad_;
|
||||
};
|
||||
|
||||
} // namespace contrib
|
||||
|
|
|
|||
|
|
@ -69,6 +69,12 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, LogSoftmaxGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, LogSoftmaxGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, LogSoftmaxGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, SoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, SoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, LogSoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, LogSoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, LogSoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_float, BatchNormalizationGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double_double, BatchNormalizationGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormalizationGrad);
|
||||
|
|
@ -187,6 +193,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16_float, InPlaceAccumulator);
|
||||
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16, SoftmaxGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16, SoftmaxGrad_13);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16, MixedPrecisionScale);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16_float, LayerNormalizationGrad);
|
||||
|
||||
|
|
@ -281,6 +288,13 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, LogSoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, LogSoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, LogSoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxGrad_13)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, SoftmaxGrad_13)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, SoftmaxGrad_13)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, LogSoftmaxGrad_13)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, LogSoftmaxGrad_13)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, LogSoftmaxGrad_13)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, 12, MLFloat16, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, 12, float, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 13, MLFloat16, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
|
|
@ -397,6 +411,7 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16_float, InPlaceAccumulator)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16, SoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16, SoftmaxGrad_13)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16, MixedPrecisionScale)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BFloat16_float, LayerNormalizationGrad)>,
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@
|
|||
#include "core/providers/cuda/cudnn_common.h"
|
||||
#include "core/providers/cuda/math/softmax.h"
|
||||
#include "core/providers/cuda/shared_inc/accumulation_type.h"
|
||||
#include "core/providers/cuda/tensor/transpose.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace cuda {
|
||||
|
|
@ -98,6 +99,16 @@ SPECIALIZED_SOFTMAXGRAD_HELPER_IMPL_BFloat16(true)
|
|||
kCudaExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
SoftmaxGrad<T>); \
|
||||
\
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
SoftmaxGrad_13, \
|
||||
kMSDomain, \
|
||||
1, \
|
||||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
SoftmaxGrad<T>); \
|
||||
\
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
LogSoftmaxGrad, \
|
||||
kMSDomain, \
|
||||
|
|
@ -105,6 +116,15 @@ SPECIALIZED_SOFTMAXGRAD_HELPER_IMPL_BFloat16(true)
|
|||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
SoftmaxGrad<T>); \
|
||||
\
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
LogSoftmaxGrad_13, \
|
||||
kMSDomain, \
|
||||
1, \
|
||||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
SoftmaxGrad<T>);
|
||||
|
||||
template <typename T>
|
||||
|
|
@ -113,16 +133,75 @@ SPECIALIZED_SOFTMAXGRAD_HELPER_IMPL_BFloat16(true)
|
|||
const TensorShape& input_shape{dY->Shape()};
|
||||
const Tensor* Y = ctx->Input<Tensor>(1);
|
||||
Tensor* dX = ctx->Output(0, input_shape);
|
||||
size_t rank = input_shape.NumDimensions();
|
||||
const size_t axis = static_cast<size_t>(HandleNegativeAxis(axis_, rank));
|
||||
bool is_transpose_required = opset_ >= 13 && axis != (rank - 1);
|
||||
|
||||
const T* dY_data = dY->template Data<T>();
|
||||
const T* Y_data = Y->template Data<T>();
|
||||
T* dX_data = dX->template MutableData<T>();
|
||||
std::unique_ptr<Tensor> transposed_dY;
|
||||
std::unique_ptr<Tensor> transposed_Y;
|
||||
std::vector<int64_t> transposed_input_dims;
|
||||
std::unique_ptr<Tensor> intermediate_output; // output that the softmax implementation will write into while using transposed input
|
||||
std::vector<size_t> permutation(rank);
|
||||
|
||||
if (log_softmax_) {
|
||||
return SoftMaxGradComputeHelper<T, true>(Stream(), dY_data, input_shape, Y_data, dX_data, CudnnHandle(), axis_);
|
||||
} else {
|
||||
return SoftMaxGradComputeHelper<T, false>(Stream(), dY_data, input_shape, Y_data, dX_data, CudnnHandle(), axis_);
|
||||
if (is_transpose_required) {
|
||||
AllocatorPtr alloc;
|
||||
auto status = ctx->GetTempSpaceAllocator(&alloc);
|
||||
if (!status.IsOK())
|
||||
return status;
|
||||
|
||||
std::iota(std::begin(permutation), std::end(permutation), 0);
|
||||
|
||||
// swap the innermost dim with the dim corresponding to axis
|
||||
permutation[axis] = rank - 1;
|
||||
permutation[rank - 1] = axis;
|
||||
|
||||
transposed_input_dims.reserve(rank);
|
||||
for (auto e : permutation) {
|
||||
transposed_input_dims.push_back(input_shape[e]);
|
||||
}
|
||||
|
||||
// Allocate a temporary tensor to hold transposed input
|
||||
auto temp_input0 = Tensor::Create(Y->DataType(), TensorShape(transposed_input_dims), alloc);
|
||||
|
||||
// Perform the transpose
|
||||
ORT_RETURN_IF_ERROR(Transpose::DoTranspose(prop_,
|
||||
Stream(),
|
||||
CublasHandle(),
|
||||
permutation, *Y, *temp_input0));
|
||||
transposed_Y = std::move(temp_input0);
|
||||
auto temp_input1 = Tensor::Create(Y->DataType(), TensorShape(transposed_input_dims), alloc);
|
||||
ORT_RETURN_IF_ERROR(Transpose::DoTranspose(prop_,
|
||||
Stream(),
|
||||
CublasHandle(),
|
||||
permutation, *dY, *temp_input1));
|
||||
transposed_dY = std::move(temp_input1);
|
||||
|
||||
// Allocate memory for the intermediate output
|
||||
intermediate_output = Tensor::Create(dX->DataType(), TensorShape(transposed_input_dims), alloc);
|
||||
}
|
||||
const T* dY_data = is_transpose_required ? transposed_dY->template Data<T>() : dY->template Data<T>();
|
||||
const T* Y_data = is_transpose_required ? transposed_Y->template Data<T>() : Y->template Data<T>();
|
||||
T* dX_data = is_transpose_required ? intermediate_output->template MutableData<T>() : dX->template MutableData<T>();
|
||||
const TensorShape* compute_input_shape = is_transpose_required ? &transposed_Y->Shape() : &input_shape;
|
||||
Status status;
|
||||
if (log_softmax_) {
|
||||
status = SoftMaxGradComputeHelper<T, true>(Stream(), dY_data, *compute_input_shape, Y_data, dX_data, CudnnHandle(), is_transpose_required ? static_cast<int64_t>(rank) - 1 : axis);
|
||||
} else {
|
||||
status = SoftMaxGradComputeHelper<T, false>(Stream(), dY_data, *compute_input_shape, Y_data, dX_data, CudnnHandle(), is_transpose_required ? static_cast<int64_t>(rank) - 1 : axis);
|
||||
}
|
||||
|
||||
if (!status.IsOK()) {
|
||||
return status;
|
||||
}
|
||||
|
||||
if (is_transpose_required) {
|
||||
// Perform the transpose to get the axes back to the original ordering
|
||||
ORT_RETURN_IF_ERROR(Transpose::DoTranspose(prop_,
|
||||
Stream(),
|
||||
CublasHandle(),
|
||||
permutation, *intermediate_output, *dX));
|
||||
}
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
#define SPECIALIZED_GRADIENT(T) \
|
||||
|
|
|
|||
|
|
@ -14,9 +14,12 @@ void dispatch_softmax_backward(cudaStream_t stream, output_t* grad_input, const
|
|||
template <typename T>
|
||||
class SoftmaxGrad final : public CudaKernel {
|
||||
public:
|
||||
SoftmaxGrad(const OpKernelInfo& info) : CudaKernel{info} {
|
||||
info.GetAttrOrDefault("axis", &axis_, static_cast<int64_t>(1));
|
||||
log_softmax_ = info.GetKernelDef().OpName() == "LogSoftmaxGrad";
|
||||
SoftmaxGrad(const OpKernelInfo& info) : CudaKernel{info},
|
||||
prop_(static_cast<const CUDAExecutionProvider*>(info.GetExecutionProvider())->GetDeviceProp()) {
|
||||
const auto& node = info.node();
|
||||
opset_ = (node.OpType() == "SoftmaxGrad_13" || node.OpType() == "LogSoftmaxGrad_13") ? 13 : 1;
|
||||
axis_ = info.GetAttrOrDefault("axis", static_cast<int64_t>(opset_ < 13 ? 1 : -1));
|
||||
log_softmax_ = info.GetKernelDef().OpName() == "LogSoftmaxGrad" || info.GetKernelDef().OpName() == "LogSoftmaxGrad_13";
|
||||
}
|
||||
|
||||
Status ComputeInternal(OpKernelContext* context) const override;
|
||||
|
|
@ -24,6 +27,8 @@ class SoftmaxGrad final : public CudaKernel {
|
|||
private:
|
||||
int64_t axis_;
|
||||
bool log_softmax_;
|
||||
int opset_; // opset_ of the forward Softmax/LogSoftmax operator
|
||||
const cudaDeviceProp& prop_;
|
||||
};
|
||||
|
||||
} // namespace cuda
|
||||
|
|
|
|||
Loading…
Reference in a new issue