From 4175859a2a999bd7900a1a267ddb757a6b79fa5a Mon Sep 17 00:00:00 2001 From: Aakash Kaushik Date: Sun, 27 Sep 2020 08:06:40 +0530 Subject: [PATCH] fixed backward function --- src/mlpack/methods/ann/layer/softmin_impl.hpp | 2 +- src/mlpack/tests/main_tests/mean_shift_test.cpp | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/mlpack/methods/ann/layer/softmin_impl.hpp b/src/mlpack/methods/ann/layer/softmin_impl.hpp index 6082a31b2a..7693ca11dd 100644 --- a/src/mlpack/methods/ann/layer/softmin_impl.hpp +++ b/src/mlpack/methods/ann/layer/softmin_impl.hpp @@ -43,7 +43,7 @@ void Softmin::Backward( const arma::Mat& gy, arma::Mat& g) { - g = input % (arma::repmat(arma::sum(gy % input), input.n_rows, 1) - gy); + g = input % (gy - arma::repmat(arma::sum(gy % input), input.n_rows, 1)); } template diff --git a/src/mlpack/tests/main_tests/mean_shift_test.cpp b/src/mlpack/tests/main_tests/mean_shift_test.cpp index eea3ceb3c7..352622a738 100644 --- a/src/mlpack/tests/main_tests/mean_shift_test.cpp +++ b/src/mlpack/tests/main_tests/mean_shift_test.cpp @@ -19,8 +19,8 @@ static const std::string testName = "MeanShift"; #include #include "test_helper.hpp" -#include -#include "../test_tools.hpp" +#include "../test_catch_tools.hpp" +#include "catch.hpp" using namespace mlpack;