diff --git a/onnxruntime/core/providers/cuda/math/softmax.cc b/onnxruntime/core/providers/cuda/math/softmax.cc index 7109f386ef..d662aacd8a 100644 --- a/onnxruntime/core/providers/cuda/math/softmax.cc +++ b/onnxruntime/core/providers/cuda/math/softmax.cc @@ -72,6 +72,22 @@ SPECIALIZED_SOFTMAX_HELPER_IMPL(MLFloat16) T, \ kCudaExecutionProvider, \ KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + Softmax); \ + ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_EX( \ + LogSoftmax, \ + kOnnxDomain, \ + 1, 10, \ + T, \ + kCudaExecutionProvider, \ + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + Softmax); \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + LogSoftmax, \ + kOnnxDomain, \ + 11, \ + T, \ + kCudaExecutionProvider, \ + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ Softmax); template @@ -84,7 +100,12 @@ Status Softmax::ComputeInternal(OpKernelContext* ctx) const { if (input_shape.Size() == 0) return Status::OK(); - return SoftMaxComputeHelper(X_data, input_shape, Y_data, CudnnHandle(), axis_); + if (log_softmax_) { + return SoftMaxComputeHelper(X_data, input_shape, Y_data, CudnnHandle(), axis_); + } + else { + return SoftMaxComputeHelper(X_data, input_shape, Y_data, CudnnHandle(), axis_); + } } #define SPECIALIZED_COMPUTE(T) \ diff --git a/onnxruntime/core/providers/cuda/math/softmax.h b/onnxruntime/core/providers/cuda/math/softmax.h index 36badfc2d9..3745231dad 100644 --- a/onnxruntime/core/providers/cuda/math/softmax.h +++ b/onnxruntime/core/providers/cuda/math/softmax.h @@ -36,12 +36,14 @@ class Softmax final : public CudaKernel { public: Softmax(const OpKernelInfo& info) : CudaKernel{info} { info.GetAttrOrDefault("axis", &axis_, static_cast(1)); + log_softmax_ = info.GetKernelDef().OpName() == "LogSoftmax"; } Status ComputeInternal(OpKernelContext* context) const override; private: int64_t axis_; + bool log_softmax_; }; } // namespace cuda diff --git a/orttraining/orttraining/core/graph/gradient_builder.cc b/orttraining/orttraining/core/graph/gradient_builder.cc index 8568b9d66b..5f83592553 100644 --- a/orttraining/orttraining/core/graph/gradient_builder.cc +++ b/orttraining/orttraining/core/graph/gradient_builder.cc @@ -675,6 +675,14 @@ IMPLEMENT_GRADIENT_BUILDER(GetSoftmaxGradient) { SrcNodeAttributes())}; } +IMPLEMENT_GRADIENT_BUILDER(GetLogSoftmaxGradient) { + return std::vector{ + NodeDef(OpDef{"LogSoftmaxGrad", kMSDomain, 1}, + {GO(0), O(0)}, + {GI(0)}, + SrcNodeAttributes())}; +} + IMPLEMENT_GRADIENT_BUILDER(GetUnsqueezeGradient) { return std::vector{ NodeDef("Squeeze", diff --git a/orttraining/orttraining/core/graph/gradient_builder.h b/orttraining/orttraining/core/graph/gradient_builder.h index 819c800820..50ee1c27e9 100644 --- a/orttraining/orttraining/core/graph/gradient_builder.h +++ b/orttraining/orttraining/core/graph/gradient_builder.h @@ -38,6 +38,7 @@ DECLARE_GRADIENT_BUILDER(GetConvGradient) DECLARE_GRADIENT_BUILDER(GetUnsqueezeGradient) DECLARE_GRADIENT_BUILDER(GetSqueezeGradient) DECLARE_GRADIENT_BUILDER(GetSoftmaxGradient) +DECLARE_GRADIENT_BUILDER(GetLogSoftmaxGradient) DECLARE_GRADIENT_BUILDER(GetSoftmaxCrossEntropyGradient) DECLARE_GRADIENT_BUILDER(GetSparseSoftmaxCrossEntropyGradient) DECLARE_GRADIENT_BUILDER(GetSoftmaxCrossEntropyLossGradient) diff --git a/orttraining/orttraining/core/graph/gradient_builder_registry.cc b/orttraining/orttraining/core/graph/gradient_builder_registry.cc index 5f50887fa7..7eb39e97d0 100644 --- a/orttraining/orttraining/core/graph/gradient_builder_registry.cc +++ b/orttraining/orttraining/core/graph/gradient_builder_registry.cc @@ -69,6 +69,7 @@ void GradientBuilderRegistry::RegisterGradientBuilders() { REGISTER_GRADIENT_BUILDER("Squeeze", GetSqueezeGradient); REGISTER_GRADIENT_BUILDER("Unsqueeze", GetUnsqueezeGradient); REGISTER_GRADIENT_BUILDER("Softmax", GetSoftmaxGradient); + REGISTER_GRADIENT_BUILDER("LogSoftmax", GetLogSoftmaxGradient); REGISTER_GRADIENT_BUILDER("SoftmaxCrossEntropy", GetSoftmaxCrossEntropyGradient); REGISTER_GRADIENT_BUILDER("SparseSoftmaxCrossEntropy", GetSparseSoftmaxCrossEntropyGradient); REGISTER_GRADIENT_BUILDER("SoftmaxCrossEntropyLoss", GetSoftmaxCrossEntropyLossGradient); diff --git a/orttraining/orttraining/core/graph/training_op_defs.cc b/orttraining/orttraining/core/graph/training_op_defs.cc index 9fa87ac094..55c390f706 100644 --- a/orttraining/orttraining/core/graph/training_op_defs.cc +++ b/orttraining/orttraining/core/graph/training_op_defs.cc @@ -351,6 +351,25 @@ void RegisterTrainingOpSchemas() { "Constrain input and output types to float tensors.") .TypeAndShapeInferenceFunction(propagateShapeAndTypeFromFirstInput); + ONNX_CONTRIB_OPERATOR_SCHEMA(LogSoftmaxGrad) + .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 axis of the inputs when coerced " + "to 2D; defaults to one because the 0th axis most likely describes " + "the batch_size", + 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 03095320c3..773b4e1a8b 100644 --- a/orttraining/orttraining/test/gradient/gradient_ops_test.cc +++ b/orttraining/orttraining/test/gradient/gradient_ops_test.cc @@ -1112,11 +1112,13 @@ TEST(GradientCheckerTest, DISABLED_BatchNormalizationGrad) { } #endif -TEST(GradientCheckerTest, SoftMaxGrad) { +void GradientCheckerSoftmaxGradHelper(bool is_log_softmax) { TensorShape shape({3, 4, 5}); float max_error; GradientChecker gradient_checker; - OpDef op_def{"Softmax"}; + + const std::string op = is_log_softmax? "LogSoftmax" : "Softmax"; + OpDef op_def{op}; // default_axis { @@ -1137,6 +1139,14 @@ TEST(GradientCheckerTest, SoftMaxGrad) { } } +TEST(GradientCheckerTest, SoftMaxGrad) { + GradientCheckerSoftmaxGradHelper(false); +} + +TEST(GradientCheckerTest, LogSoftMaxGrad) { + GradientCheckerSoftmaxGradHelper(true); +} + void TestSoftmaxCrossEntropyGrad(const TensorShape& input_shape, const std::string& reduction) { float max_error; GradientChecker gradient_checker; diff --git a/orttraining/orttraining/test/training_ops/cuda/softmax_test.cc b/orttraining/orttraining/test/training_ops/cuda/softmax_test.cc index 8501479877..192c7cb2b7 100644 --- a/orttraining/orttraining/test/training_ops/cuda/softmax_test.cc +++ b/orttraining/orttraining/test/training_ops/cuda/softmax_test.cc @@ -8,9 +8,12 @@ namespace test { static void TestSoftmax(const std::vector& X_dims, const std::vector& Y_dims, + bool is_log_softmax=false, double per_sample_tolerance = 1e-4, double relative_per_sample_tolerance = 1e-4) { - CompareOpTester test("Softmax"); + + const char* op = is_log_softmax? "LogSoftmax" : "Softmax"; + CompareOpTester test(op); // create rand inputs RandomValueGenerator random{}; @@ -26,21 +29,36 @@ static void TestSoftmax(const std::vector& X_dims, TEST(CudaKernelTest, Softmax_SmallTensor) { std::vector X_dims{8, 2, 128, 128}; std::vector Y_dims{8, 2, 128, 128}; - TestSoftmax(X_dims, Y_dims); + TestSoftmax(X_dims, Y_dims, false); } TEST(CudaKernelTest, Softmax_LargeTensor) { std::vector X_dims{8, 16, 512, 512}; std::vector Y_dims{8, 16, 512, 512}; - TestSoftmax(X_dims, Y_dims); + TestSoftmax(X_dims, Y_dims, false); +} + +TEST(CudaKernelTest, LogSoftmax_SmallTensor) { + std::vector X_dims{8, 2, 128, 128}; + std::vector Y_dims{8, 2, 128, 128}; + TestSoftmax(X_dims, Y_dims, true); +} + +TEST(CudaKernelTest, LogSoftmax_LargeTensor) { + std::vector X_dims{8, 16, 512, 512}; + std::vector Y_dims{8, 16, 512, 512}; + TestSoftmax(X_dims, Y_dims, true); } static void TestSoftmaxGrad(const std::vector& dY_dims, const std::vector& Y_dims, const std::vector& dX_dims, + bool is_log_softmax = false, double per_sample_tolerance = 1e-4, - double relative_per_sample_tolerance = 1e-4) { - CompareOpTester test("SoftmaxGrad", 1, kMSDomain); + double relative_per_sample_tolerance = 1e-4) { + + const char* op = is_log_softmax? "LogSoftmaxGrad" : "SoftmaxGrad"; + CompareOpTester test(op, 1, kMSDomain); // create rand inputs RandomValueGenerator random{}; @@ -71,5 +89,19 @@ TEST(CudaKernelTest, SoftmaxGrad_LargeTensor) { TestSoftmaxGrad(dY_dims, Y_dims, dX_dims); } +TEST(CudaKernelTest, LogSoftmaxGrad_SmallTensor) { + std::vector dY_dims{8, 2, 128, 128}; + std::vector Y_dims{8, 2, 128, 128}; + std::vector dX_dims{8, 2, 128, 128}; + TestSoftmaxGrad(dY_dims, Y_dims, dX_dims, true); +} + +TEST(CudaKernelTest, LogSoftmaxGrad_LargeTensor) { + std::vector dY_dims{8, 16, 512, 512}; + std::vector Y_dims{8, 16, 512, 512}; + std::vector dX_dims{8, 16, 512, 512}; + TestSoftmaxGrad(dY_dims, Y_dims, dX_dims, true); +} + } // namespace test } // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc index ab37511dc4..67d753aae6 100644 --- a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc @@ -36,6 +36,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sin class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, ConvGrad); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 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, kOnnxDomain, 9, AveragePoolGrad); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MaxPoolGrad); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherGrad); @@ -127,6 +128,7 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) { 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 0ffcaf9ca4..4ae95b7a31 100644 --- a/orttraining/orttraining/training_ops/cpu/op_gradients.cc +++ b/orttraining/orttraining/training_ops/cpu/op_gradients.cc @@ -100,5 +100,51 @@ Status SoftmaxGrad::Compute(OpKernelContext* context) const { return Status::OK(); } +ONNX_OPERATOR_KERNEL_EX( + LogSoftmaxGrad, + kMSDomain, + 1, + kCpuExecutionProvider, + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), + LogSoftmaxGrad); + +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(); +} + } // namespace contrib } // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cpu/op_gradients.h b/orttraining/orttraining/training_ops/cpu/op_gradients.h index 7fa4697ab4..22da17a446 100644 --- a/orttraining/orttraining/training_ops/cpu/op_gradients.h +++ b/orttraining/orttraining/training_ops/cpu/op_gradients.h @@ -46,5 +46,20 @@ class SoftmaxGrad final : public OpKernel { 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_; +}; + } // namespace contrib } // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc index cbbb7dd307..7934c9d991 100644 --- a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc @@ -56,6 +56,9 @@ class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomai class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, SoftmaxGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, SoftmaxGrad); +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, BatchNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, BatchNormalizationGrad); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, GatherGrad); @@ -183,6 +186,9 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + 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 7c1fef3f5d..615ca9792c 100644 --- a/orttraining/orttraining/training_ops/cuda/math/softmax_grad.cc +++ b/orttraining/orttraining/training_ops/cuda/math/softmax_grad.cc @@ -9,37 +9,28 @@ namespace onnxruntime { namespace cuda { -#define REGISTER_GRADIENT_KERNEL_TYPED(T) \ - ONNX_OPERATOR_TYPED_KERNEL_EX( \ - SoftmaxGrad, \ - kMSDomain, \ - 1, \ - T, \ - kCudaExecutionProvider, \ - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ - SoftmaxGrad); - -template -Status SoftmaxGrad::ComputeInternal(OpKernelContext* ctx) const { +template +Status SoftMaxGradComputeHelper( + const T* dY, + const TensorShape& input_shape, + const T* Y, + T* dX, + cudnnHandle_t handle, + int64_t axis) { typedef typename ToCudaType::MappedType CudaT; - const Tensor* dY = ctx->Input(0); - const TensorShape input_shape{dY->Shape()}; - - const Tensor* Y = ctx->Input(1); - - const int64_t normalized_axis = HandleNegativeAxis(axis_, input_shape.NumDimensions()); + const int64_t normalized_axis = HandleNegativeAxis(axis, input_shape.NumDimensions()); int64_t N = input_shape.SizeToDimension(normalized_axis); int64_t D = input_shape.SizeFromDimension(normalized_axis); std::vector dims({N, 1, 1, D}); // cudnn expects 4D shape in NCHW format - auto dY_data = reinterpret_cast(dY->template Data()); - auto Y_data = reinterpret_cast(Y->template Data()); - auto dX_data = reinterpret_cast(ctx->Output(0, input_shape)->template MutableData()); + auto dY_data = reinterpret_cast(dY); + auto Y_data = reinterpret_cast(Y); + auto dX_data = reinterpret_cast(dX); if (D == input_shape[normalized_axis] && D <= 1024 && D * sizeof(T) <= 4096) { - dispatch_softmax_backward, false>(dX_data, dY_data, Y_data, gsl::narrow_cast(D), gsl::narrow_cast(D), gsl::narrow_cast(N)); + dispatch_softmax_backward, is_log_softmax>(dX_data, dY_data, Y_data, gsl::narrow_cast(D), gsl::narrow_cast(D), gsl::narrow_cast(N)); return Status::OK(); } @@ -51,8 +42,8 @@ Status SoftmaxGrad::ComputeInternal(OpKernelContext* ctx) const { ORT_RETURN_IF_ERROR(output_tensor.Set(dims, CudnnTensor::GetDataType())); CUDNN_RETURN_IF_ERROR( cudnnSoftmaxBackward( - CudnnHandle(), - CUDNN_SOFTMAX_ACCURATE, + handle, + is_log_softmax? CUDNN_SOFTMAX_LOG : CUDNN_SOFTMAX_ACCURATE, CUDNN_SOFTMAX_MODE_INSTANCE, &alpha, input_tensor, @@ -66,6 +57,45 @@ Status SoftmaxGrad::ComputeInternal(OpKernelContext* ctx) const { return Status::OK(); } + +#define REGISTER_GRADIENT_KERNEL_TYPED(T) \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + SoftmaxGrad, \ + kMSDomain, \ + 1, \ + T, \ + kCudaExecutionProvider, \ + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + SoftmaxGrad); \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + LogSoftmaxGrad, \ + kMSDomain, \ + 1, \ + T, \ + kCudaExecutionProvider, \ + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + SoftmaxGrad); + +template +Status SoftmaxGrad::ComputeInternal(OpKernelContext* ctx) const { + + const Tensor* dY = ctx->Input(0); + const TensorShape& input_shape{dY->Shape()}; + const Tensor* Y = ctx->Input(1); + Tensor* dX = ctx->Output(0, input_shape); + + const T* dY_data = dY->template Data(); + const T* Y_data = Y->template Data(); + T* dX_data = dX->template MutableData(); + + if (log_softmax_) { + return SoftMaxGradComputeHelper(dY_data, input_shape, Y_data, dX_data, CudnnHandle(), axis_); + } + else { + return SoftMaxGradComputeHelper(dY_data, input_shape, Y_data, dX_data, CudnnHandle(), axis_); + } +} + #define SPECIALIZED_GRADIENT(T) \ REGISTER_GRADIENT_KERNEL_TYPED(T) \ template Status SoftmaxGrad::ComputeInternal(OpKernelContext* ctx) const; diff --git a/orttraining/orttraining/training_ops/cuda/math/softmax_grad.h b/orttraining/orttraining/training_ops/cuda/math/softmax_grad.h index 063d20ba36..4d2e6bae4a 100644 --- a/orttraining/orttraining/training_ops/cuda/math/softmax_grad.h +++ b/orttraining/orttraining/training_ops/cuda/math/softmax_grad.h @@ -16,12 +16,14 @@ 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"; } Status ComputeInternal(OpKernelContext* context) const override; private: int64_t axis_; + bool log_softmax_; }; } // namespace cuda diff --git a/orttraining/orttraining/training_ops/cuda/math/softmax_grad_impl.cu b/orttraining/orttraining/training_ops/cuda/math/softmax_grad_impl.cu index 0f4140c1f8..294f955ac1 100644 --- a/orttraining/orttraining/training_ops/cuda/math/softmax_grad_impl.cu +++ b/orttraining/orttraining/training_ops/cuda/math/softmax_grad_impl.cu @@ -181,7 +181,8 @@ void dispatch_softmax_backward(output_t* grad_input, const input_t* grad, const } #define SPECIALIZED_SOFTMAX_GRAD_IMPL(input_t, output_t, acc_t) \ -template void dispatch_softmax_backward(input_t * grad_input, const output_t* grad, const output_t* output, int softmax_elements, int softmax_elements_stride, int batch_count); +template void dispatch_softmax_backward(input_t * grad_input, const output_t* grad, const output_t* output, int softmax_elements, int softmax_elements_stride, int batch_count); \ +template void dispatch_softmax_backward(input_t * grad_input, const output_t* grad, const output_t* output, int softmax_elements, int softmax_elements_stride, int batch_count); SPECIALIZED_SOFTMAX_GRAD_IMPL(float, float, float) SPECIALIZED_SOFTMAX_GRAD_IMPL(half, half, float)