diff --git a/src/mlpack/methods/decision_tree/mad_gain.hpp b/src/mlpack/methods/decision_tree/mad_gain.hpp index 065de37dcb..1f3e460a77 100644 --- a/src/mlpack/methods/decision_tree/mad_gain.hpp +++ b/src/mlpack/methods/decision_tree/mad_gain.hpp @@ -109,6 +109,8 @@ class MADGain accWeights[0] += accWeights[1] + accWeights[2] + accWeights[3]; weightedMean[0] += weightedMean[1] + weightedMean[2] + weightedMean[3]; + weightedMean[0] /= (double) (end - begin); + std::cout << "WeightedMean: " << weightedMean[0] << std::endl; // Catch edge case: if there are no weights, the impurity is zero. if (accWeights[0] == 0.0) @@ -152,6 +154,8 @@ class MADGain } mean[0] += mean[1] + mean[2] + mean[3]; + mean[0] /= (double) (end - begin); + std::cout << "Mean: " << mean[0] << std::endl; for (size_t i = begin; i < end; ++i) mad += std::abs(labels[i] - mean[0]); diff --git a/src/mlpack/tests/decision_tree_test.cpp b/src/mlpack/tests/decision_tree_test.cpp index 8aaa7194e2..ab8aa094ae 100644 --- a/src/mlpack/tests/decision_tree_test.cpp +++ b/src/mlpack/tests/decision_tree_test.cpp @@ -39,13 +39,13 @@ TEST_CASE("MADGainPerfectTest", "[DecisionTreeRegressionTest]") } /** - * Make sure that for a normal distribution of labels, - * MAD_gain = mean of absolute values of the distribution. + * Make sure that when mean of labels is zero, MAD_gain = mean of + * absolute values of the distribution. */ TEST_CASE("MADGainNormalTest", "[DecisionTreeRegressionTest") { arma::rowvec weights(10, arma::fill::ones); - arma::rowvec labels(10, arma::fill::randn); // Mean = 0. + arma::rowvec labels = { 1, 2, 3, 4, 5, -1, -2, -3, -4, -5 }; // Mean = 0. // Theoretical gain. double theoreticalGain = 0.0;