Updated SoftmaxGrad and LogSoftmaxGrad to support version 13. (#9733)

* Updated SoftmaxGrad_13/LogSoftmaxGrad_13 to support version 13.
This commit is contained in:
satyajandhyala 2021-11-18 17:39:16 -08:00 committed by GitHub
parent dc1724b0e2
commit 3af14fc554
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
10 changed files with 269 additions and 85 deletions

View file

@ -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

View file

@ -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())};

View file

@ -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")

View file

@ -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) {

View file

@ -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)>,

View file

@ -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,

View file

@ -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

View file

@ -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)>,

View file

@ -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) \

View file

@ -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