diff --git a/onnxruntime/test/testdata/kernel_def_hashes/training_ops.cpu.json b/onnxruntime/test/testdata/kernel_def_hashes/training_ops.cpu.json index 764b882112..856996d074 100644 --- a/onnxruntime/test/testdata/kernel_def_hashes/training_ops.cpu.json +++ b/onnxruntime/test/testdata/kernel_def_hashes/training_ops.cpu.json @@ -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 diff --git a/orttraining/orttraining/core/graph/gradient_builder.cc b/orttraining/orttraining/core/graph/gradient_builder.cc index 1fad171dae..342132e5c3 100755 --- a/orttraining/orttraining/core/graph/gradient_builder.cc +++ b/orttraining/orttraining/core/graph/gradient_builder.cc @@ -707,7 +707,7 @@ IMPLEMENT_GRADIENT_BUILDER(GetSigmoidGradient) { IMPLEMENT_GRADIENT_BUILDER(GetSoftmaxGradient) { return std::vector{ - 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(OpDef{"LogSoftmaxGrad", kMSDomain, 1}, + NodeDef(OpDef{SrcNodeOpsetVersion() < 13 ? "LogSoftmaxGrad" : "LogSoftmaxGrad_13", kMSDomain, 1}, {GO(0), O(0)}, {GI(0)}, SrcNodeAttributes())}; diff --git a/orttraining/orttraining/core/graph/training_op_defs.cc b/orttraining/orttraining/core/graph/training_op_defs.cc index e14f0f2d9e..749032ea9d 100644 --- a/orttraining/orttraining/core/graph/training_op_defs.cc +++ b/orttraining/orttraining/core/graph/training_op_defs.cc @@ -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(-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(-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") diff --git a/orttraining/orttraining/test/gradient/gradient_ops_test.cc b/orttraining/orttraining/test/gradient/gradient_ops_test.cc index 53f885d8eb..4535cd5e3f 100644 --- a/orttraining/orttraining/test/gradient/gradient_ops_test.cc +++ b/orttraining/orttraining/test/gradient/gradient_ops_test.cc @@ -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 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) { diff --git a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc index 76d4f9dde5..800fd48221 100644 --- a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc @@ -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, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/cpu/op_gradients.cc b/orttraining/orttraining/training_ops/cpu/op_gradients.cc index bbcfdc6ec8..a5ef415374 100644 --- a/orttraining/orttraining/training_ops/cpu/op_gradients.cc +++ b/orttraining/orttraining/training_ops/cpu/op_gradients.cc @@ -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()), SoftmaxGrad); +ONNX_OPERATOR_KERNEL_EX( + SoftmaxGrad_13, + kMSDomain, + 1, + kCpuExecutionProvider, + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), + SoftmaxGrad); + template Status SoftmaxGrad::Compute(OpKernelContext* context) const { auto& dY = *context->Input(0); @@ -65,7 +74,8 @@ Status SoftmaxGrad::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(HandleNegativeAxis(axis_, rank)); size_t N = input_shape.SizeToDimension(axis); size_t D = input_shape.SizeFromDimension(axis); @@ -74,30 +84,88 @@ Status SoftmaxGrad::Compute(OpKernelContext* context) const { return Status::OK(); } - std::vector scale_(N); - std::vector sum_multiplier_(D, 1.f); // initialize all multiplier values to 1.0 + bool is_transpose_required = opset_ >= 13 && axis != (rank - 1); + + std::unique_ptr transposed_dY; + std::unique_ptr transposed_Y; + std::vector transposed_input_dims; + std::unique_ptr intermediate_output; // output that the softmax implementation will write into while using transposed input + std::vector 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(N); const int d = gsl::narrow_cast(D); const int nd = gsl::narrow_cast(N * D); - - float* scaledata = scale_.data(); - const float* Ydata = Y.template Data(); - const float* dYdata = dY.template Data(); - float* dXdata = dX.template MutableData(); + const float* Ydata = is_transpose_required ? transposed_Y->template Data() : Y.template Data(); + const float* dYdata = is_transpose_required ? transposed_dY->template Data() : dY.template Data(); + float* dXdata = is_transpose_required ? intermediate_output->template MutableData() : dX.template MutableData(); gsl::copy(gsl::make_span(dYdata, nd), gsl::make_span(dXdata, nd)); + if (is_logsoftmaxgrad_) { + std::vector eY(nd); + float* eYdata = eY.data(); - for (size_t i = 0; i < N; ++i) { - math::Dot(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(nd, Ydata, eYdata, nullptr); + for (size_t i = 0; i < N; ++i) { + float sdY; + math::Sum(d, dYdata + i * d, &sdY, nullptr, nullptr); + math::Axpy(d, -sdY, eYdata + i * d, dXdata + i * d, nullptr); + } + } else { + std::vector scale_(N); + std::vector 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(d, Ydata + i * d, dYdata + i * d, + scaledata + i, nullptr); + } + + concurrency::ThreadPool* tp = context->GetOperatorThreadPool(); + math::Gemm(CblasNoTrans, CblasNoTrans, n, d, 1, -1, + scaledata, sum_multiplier_.data(), 1, + dXdata, tp); + + math::Mul(gsl::narrow_cast(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(CblasNoTrans, CblasNoTrans, n, d, 1, -1, - scaledata, sum_multiplier_.data(), 1, - dXdata, tp); - - math::Mul(gsl::narrow_cast(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()), - LogSoftmaxGrad); + SoftmaxGrad); -template -Status LogSoftmaxGrad::Compute(OpKernelContext* context) const { - auto& dY = *context->Input(0); - auto& Y = *context->Input(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(D); - const int nd = gsl::narrow_cast(N * D); - - const float* Ydata = Y.template Data(); - const float* dYdata = dY.template Data(); - float* dXdata = dX.template MutableData(); - - std::vector 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(nd, Ydata, eYdata, nullptr); - for (size_t i = 0; i < N; ++i) { - float sdY; - math::Sum(d, dYdata + i * d, &sdY, nullptr, nullptr); - math::Axpy(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()), + SoftmaxGrad); ONNX_OPERATOR_KERNEL_EX( SigmoidGrad, diff --git a/orttraining/orttraining/training_ops/cpu/op_gradients.h b/orttraining/orttraining/training_ops/cpu/op_gradients.h index 80081c15c4..4e519a7622 100644 --- a/orttraining/orttraining/training_ops/cpu/op_gradients.h +++ b/orttraining/orttraining/training_ops/cpu/op_gradients.h @@ -61,7 +61,10 @@ template class SoftmaxGrad final : public OpKernel { public: explicit SoftmaxGrad(const OpKernelInfo& info) : OpKernel(info) { - axis_ = info.GetAttrOrDefault("axis", 0); + const auto& node = info.node(); + opset_ = (node.OpType() == "SoftmaxGrad_13" || node.OpType() == "LogSoftmaxGrad_13") ? 13 : 1; + axis_ = info.GetAttrOrDefault("axis", static_cast(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 -class LogSoftmaxGrad final : public OpKernel { - public: - explicit LogSoftmaxGrad(const OpKernelInfo& info) : OpKernel(info) { - axis_ = info.GetAttrOrDefault("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 diff --git a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc index 394b8ad7c2..8d0882f2ee 100644 --- a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc @@ -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, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -397,6 +411,7 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/cuda/math/softmax_grad.cc b/orttraining/orttraining/training_ops/cuda/math/softmax_grad.cc index 9e00f9ede1..7ce8f0184a 100644 --- a/orttraining/orttraining/training_ops/cuda/math/softmax_grad.cc +++ b/orttraining/orttraining/training_ops/cuda/math/softmax_grad.cc @@ -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()), \ SoftmaxGrad); \ + \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + SoftmaxGrad_13, \ + kMSDomain, \ + 1, \ + T, \ + kCudaExecutionProvider, \ + (*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + SoftmaxGrad); \ + \ 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()), \ + SoftmaxGrad); \ + \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + LogSoftmaxGrad_13, \ + kMSDomain, \ + 1, \ + T, \ + kCudaExecutionProvider, \ + (*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType()), \ SoftmaxGrad); template @@ -113,16 +133,75 @@ SPECIALIZED_SOFTMAXGRAD_HELPER_IMPL_BFloat16(true) const TensorShape& input_shape{dY->Shape()}; const Tensor* Y = ctx->Input(1); Tensor* dX = ctx->Output(0, input_shape); + size_t rank = input_shape.NumDimensions(); + const size_t axis = static_cast(HandleNegativeAxis(axis_, rank)); + bool is_transpose_required = opset_ >= 13 && axis != (rank - 1); - const T* dY_data = dY->template Data(); - const T* Y_data = Y->template Data(); - T* dX_data = dX->template MutableData(); + std::unique_ptr transposed_dY; + std::unique_ptr transposed_Y; + std::vector transposed_input_dims; + std::unique_ptr intermediate_output; // output that the softmax implementation will write into while using transposed input + std::vector permutation(rank); - if (log_softmax_) { - return SoftMaxGradComputeHelper(Stream(), dY_data, input_shape, Y_data, dX_data, CudnnHandle(), axis_); - } else { - return SoftMaxGradComputeHelper(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() : dY->template Data(); + const T* Y_data = is_transpose_required ? transposed_Y->template Data() : Y->template Data(); + T* dX_data = is_transpose_required ? intermediate_output->template MutableData() : dX->template MutableData(); + const TensorShape* compute_input_shape = is_transpose_required ? &transposed_Y->Shape() : &input_shape; + Status status; + if (log_softmax_) { + status = SoftMaxGradComputeHelper(Stream(), dY_data, *compute_input_shape, Y_data, dX_data, CudnnHandle(), is_transpose_required ? static_cast(rank) - 1 : axis); + } else { + status = SoftMaxGradComputeHelper(Stream(), dY_data, *compute_input_shape, Y_data, dX_data, CudnnHandle(), is_transpose_required ? static_cast(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) \ diff --git a/orttraining/orttraining/training_ops/cuda/math/softmax_grad.h b/orttraining/orttraining/training_ops/cuda/math/softmax_grad.h index 4e50cf2cf4..21543ca9a2 100644 --- a/orttraining/orttraining/training_ops/cuda/math/softmax_grad.h +++ b/orttraining/orttraining/training_ops/cuda/math/softmax_grad.h @@ -14,9 +14,12 @@ void dispatch_softmax_backward(cudaStream_t stream, output_t* grad_input, const template class SoftmaxGrad final : public CudaKernel { public: - SoftmaxGrad(const OpKernelInfo& info) : CudaKernel{info} { - info.GetAttrOrDefault("axis", &axis_, static_cast(1)); - log_softmax_ = info.GetKernelDef().OpName() == "LogSoftmaxGrad"; + SoftmaxGrad(const OpKernelInfo& info) : CudaKernel{info}, + prop_(static_cast(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(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