mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-28 20:11:22 +00:00
Wire log(softmax) grad cuda kernel and add log(softmax) grad cpu kernel (#4726)
* logsoftmax cuda kernel * add cpu logsoftmaxgrad * revert debug printout * revert disable for debug builds * use /alpha x + y instead * remove misleading log_softmax_ bool Co-authored-by: suffian khan <sukha@OrtTrainingDev1.af05slrtruoetgaxwwjv5nsq5e.px.internal.cloudapp.net>
This commit is contained in:
parent
9a73c8f448
commit
4d39c6a6cb
15 changed files with 229 additions and 33 deletions
|
|
@ -72,6 +72,22 @@ SPECIALIZED_SOFTMAX_HELPER_IMPL(MLFloat16)
|
|||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
Softmax<T>); \
|
||||
ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_EX( \
|
||||
LogSoftmax, \
|
||||
kOnnxDomain, \
|
||||
1, 10, \
|
||||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
Softmax<T>); \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
LogSoftmax, \
|
||||
kOnnxDomain, \
|
||||
11, \
|
||||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
Softmax<T>);
|
||||
|
||||
template <typename T>
|
||||
|
|
@ -84,7 +100,12 @@ Status Softmax<T>::ComputeInternal(OpKernelContext* ctx) const {
|
|||
if (input_shape.Size() == 0)
|
||||
return Status::OK();
|
||||
|
||||
return SoftMaxComputeHelper<T, false>(X_data, input_shape, Y_data, CudnnHandle(), axis_);
|
||||
if (log_softmax_) {
|
||||
return SoftMaxComputeHelper<T, true>(X_data, input_shape, Y_data, CudnnHandle(), axis_);
|
||||
}
|
||||
else {
|
||||
return SoftMaxComputeHelper<T, false>(X_data, input_shape, Y_data, CudnnHandle(), axis_);
|
||||
}
|
||||
}
|
||||
|
||||
#define SPECIALIZED_COMPUTE(T) \
|
||||
|
|
|
|||
|
|
@ -36,12 +36,14 @@ class Softmax final : public CudaKernel {
|
|||
public:
|
||||
Softmax(const OpKernelInfo& info) : CudaKernel{info} {
|
||||
info.GetAttrOrDefault("axis", &axis_, static_cast<int64_t>(1));
|
||||
log_softmax_ = info.GetKernelDef().OpName() == "LogSoftmax";
|
||||
}
|
||||
|
||||
Status ComputeInternal(OpKernelContext* context) const override;
|
||||
|
||||
private:
|
||||
int64_t axis_;
|
||||
bool log_softmax_;
|
||||
};
|
||||
|
||||
} // namespace cuda
|
||||
|
|
|
|||
|
|
@ -675,6 +675,14 @@ IMPLEMENT_GRADIENT_BUILDER(GetSoftmaxGradient) {
|
|||
SrcNodeAttributes())};
|
||||
}
|
||||
|
||||
IMPLEMENT_GRADIENT_BUILDER(GetLogSoftmaxGradient) {
|
||||
return std::vector<NodeDef>{
|
||||
NodeDef(OpDef{"LogSoftmaxGrad", kMSDomain, 1},
|
||||
{GO(0), O(0)},
|
||||
{GI(0)},
|
||||
SrcNodeAttributes())};
|
||||
}
|
||||
|
||||
IMPLEMENT_GRADIENT_BUILDER(GetUnsqueezeGradient) {
|
||||
return std::vector<NodeDef>{
|
||||
NodeDef("Squeeze",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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<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")
|
||||
|
|
|
|||
|
|
@ -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<float, float, float> 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<float, float, float> gradient_checker;
|
||||
|
|
|
|||
|
|
@ -8,9 +8,12 @@ namespace test {
|
|||
|
||||
static void TestSoftmax(const std::vector<int64_t>& X_dims,
|
||||
const std::vector<int64_t>& 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<int64_t>& X_dims,
|
|||
TEST(CudaKernelTest, Softmax_SmallTensor) {
|
||||
std::vector<int64_t> X_dims{8, 2, 128, 128};
|
||||
std::vector<int64_t> Y_dims{8, 2, 128, 128};
|
||||
TestSoftmax(X_dims, Y_dims);
|
||||
TestSoftmax(X_dims, Y_dims, false);
|
||||
}
|
||||
|
||||
TEST(CudaKernelTest, Softmax_LargeTensor) {
|
||||
std::vector<int64_t> X_dims{8, 16, 512, 512};
|
||||
std::vector<int64_t> Y_dims{8, 16, 512, 512};
|
||||
TestSoftmax(X_dims, Y_dims);
|
||||
TestSoftmax(X_dims, Y_dims, false);
|
||||
}
|
||||
|
||||
TEST(CudaKernelTest, LogSoftmax_SmallTensor) {
|
||||
std::vector<int64_t> X_dims{8, 2, 128, 128};
|
||||
std::vector<int64_t> Y_dims{8, 2, 128, 128};
|
||||
TestSoftmax(X_dims, Y_dims, true);
|
||||
}
|
||||
|
||||
TEST(CudaKernelTest, LogSoftmax_LargeTensor) {
|
||||
std::vector<int64_t> X_dims{8, 16, 512, 512};
|
||||
std::vector<int64_t> Y_dims{8, 16, 512, 512};
|
||||
TestSoftmax(X_dims, Y_dims, true);
|
||||
}
|
||||
|
||||
static void TestSoftmaxGrad(const std::vector<int64_t>& dY_dims,
|
||||
const std::vector<int64_t>& Y_dims,
|
||||
const std::vector<int64_t>& 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<int64_t> dY_dims{8, 2, 128, 128};
|
||||
std::vector<int64_t> Y_dims{8, 2, 128, 128};
|
||||
std::vector<int64_t> dX_dims{8, 2, 128, 128};
|
||||
TestSoftmaxGrad(dY_dims, Y_dims, dX_dims, true);
|
||||
}
|
||||
|
||||
TEST(CudaKernelTest, LogSoftmaxGrad_LargeTensor) {
|
||||
std::vector<int64_t> dY_dims{8, 16, 512, 512};
|
||||
std::vector<int64_t> Y_dims{8, 16, 512, 512};
|
||||
std::vector<int64_t> dX_dims{8, 16, 512, 512};
|
||||
TestSoftmaxGrad(dY_dims, Y_dims, dX_dims, true);
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -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<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, ConvGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 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, kOnnxDomain, 9, AveragePoolGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MaxPoolGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherGrad)>,
|
||||
|
|
|
|||
|
|
@ -100,5 +100,51 @@ Status SoftmaxGrad<T>::Compute(OpKernelContext* context) const {
|
|||
return Status::OK();
|
||||
}
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
LogSoftmaxGrad,
|
||||
kMSDomain,
|
||||
1,
|
||||
kCpuExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
LogSoftmaxGrad<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();
|
||||
}
|
||||
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -46,5 +46,20 @@ class SoftmaxGrad final : public OpKernel {
|
|||
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_;
|
||||
};
|
||||
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -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<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, SoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, SoftmaxGrad)>,
|
||||
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_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, int64_t, SoftmaxCrossEntropyLossGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, BatchNormalizationGrad)>,
|
||||
|
|
|
|||
|
|
@ -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<T>()), \
|
||||
SoftmaxGrad<T>);
|
||||
|
||||
template <typename T>
|
||||
Status SoftmaxGrad<T>::ComputeInternal(OpKernelContext* ctx) const {
|
||||
template <typename T, bool is_log_softmax>
|
||||
Status SoftMaxGradComputeHelper(
|
||||
const T* dY,
|
||||
const TensorShape& input_shape,
|
||||
const T* Y,
|
||||
T* dX,
|
||||
cudnnHandle_t handle,
|
||||
int64_t axis) {
|
||||
typedef typename ToCudaType<T>::MappedType CudaT;
|
||||
|
||||
const Tensor* dY = ctx->Input<Tensor>(0);
|
||||
const TensorShape input_shape{dY->Shape()};
|
||||
|
||||
const Tensor* Y = ctx->Input<Tensor>(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<int64_t> dims({N, 1, 1, D}); // cudnn expects 4D shape in NCHW format
|
||||
|
||||
auto dY_data = reinterpret_cast<const CudaT*>(dY->template Data<T>());
|
||||
auto Y_data = reinterpret_cast<const CudaT*>(Y->template Data<T>());
|
||||
auto dX_data = reinterpret_cast<CudaT*>(ctx->Output(0, input_shape)->template MutableData<T>());
|
||||
auto dY_data = reinterpret_cast<const CudaT*>(dY);
|
||||
auto Y_data = reinterpret_cast<const CudaT*>(Y);
|
||||
auto dX_data = reinterpret_cast<CudaT*>(dX);
|
||||
|
||||
if (D == input_shape[normalized_axis] && D <= 1024 && D * sizeof(T) <= 4096) {
|
||||
dispatch_softmax_backward<CudaT, CudaT, AccType<T>, false>(dX_data, dY_data, Y_data, gsl::narrow_cast<int>(D), gsl::narrow_cast<int>(D), gsl::narrow_cast<int>(N));
|
||||
dispatch_softmax_backward<CudaT, CudaT, AccType<T>, is_log_softmax>(dX_data, dY_data, Y_data, gsl::narrow_cast<int>(D), gsl::narrow_cast<int>(D), gsl::narrow_cast<int>(N));
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
|
|
@ -51,8 +42,8 @@ Status SoftmaxGrad<T>::ComputeInternal(OpKernelContext* ctx) const {
|
|||
ORT_RETURN_IF_ERROR(output_tensor.Set(dims, CudnnTensor::GetDataType<CudaT>()));
|
||||
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<T>::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<T>()), \
|
||||
SoftmaxGrad<T>); \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
LogSoftmaxGrad, \
|
||||
kMSDomain, \
|
||||
1, \
|
||||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
SoftmaxGrad<T>);
|
||||
|
||||
template <typename T>
|
||||
Status SoftmaxGrad<T>::ComputeInternal(OpKernelContext* ctx) const {
|
||||
|
||||
const Tensor* dY = ctx->Input<Tensor>(0);
|
||||
const TensorShape& input_shape{dY->Shape()};
|
||||
const Tensor* Y = ctx->Input<Tensor>(1);
|
||||
Tensor* dX = ctx->Output(0, input_shape);
|
||||
|
||||
const T* dY_data = dY->template Data<T>();
|
||||
const T* Y_data = Y->template Data<T>();
|
||||
T* dX_data = dX->template MutableData<T>();
|
||||
|
||||
if (log_softmax_) {
|
||||
return SoftMaxGradComputeHelper<T, true>(dY_data, input_shape, Y_data, dX_data, CudnnHandle(), axis_);
|
||||
}
|
||||
else {
|
||||
return SoftMaxGradComputeHelper<T, false>(dY_data, input_shape, Y_data, dX_data, CudnnHandle(), axis_);
|
||||
}
|
||||
}
|
||||
|
||||
#define SPECIALIZED_GRADIENT(T) \
|
||||
REGISTER_GRADIENT_KERNEL_TYPED(T) \
|
||||
template Status SoftmaxGrad<T>::ComputeInternal(OpKernelContext* ctx) const;
|
||||
|
|
|
|||
|
|
@ -16,12 +16,14 @@ 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";
|
||||
}
|
||||
|
||||
Status ComputeInternal(OpKernelContext* context) const override;
|
||||
|
||||
private:
|
||||
int64_t axis_;
|
||||
bool log_softmax_;
|
||||
};
|
||||
|
||||
} // namespace cuda
|
||||
|
|
|
|||
|
|
@ -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, output_t, acc_t, false>(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, output_t, acc_t, false>(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, output_t, acc_t, true>(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)
|
||||
|
|
|
|||
Loading…
Reference in a new issue