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:
suffiank 2020-08-06 10:49:08 -07:00 committed by GitHub
parent 9a73c8f448
commit 4d39c6a6cb
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
15 changed files with 229 additions and 33 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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