diff --git a/src/mlpack/methods/decision_tree/all_categorical_split.hpp b/src/mlpack/methods/decision_tree/all_categorical_split.hpp index faa8f16c6b..09f717c155 100644 --- a/src/mlpack/methods/decision_tree/all_categorical_split.hpp +++ b/src/mlpack/methods/decision_tree/all_categorical_split.hpp @@ -47,23 +47,23 @@ class AllCategoricalSplit * @param weights Weights associated with labels. * @param minimumLeafSize Minimum number of points in a leaf node for * splitting. - * @param classProbabilities Class probabilities vector, which may be filled - * with split information a successful split. + * @param splitInfo Stores split information on a successful split. * @param minimumGainSplit Minimum gain split. * @param aux Auxiliary split information, which may be modified on a * successful split. */ - template + template static double SplitIfBetter( const double bestGain, const VecType& data, const size_t numCategories, - const arma::Row& labels, + const arma::Row& labels, + const size_t begin, const size_t numClasses, const WeightVecType& weights, const size_t minimumLeafSize, const double minimumGainSplit, - arma::vec& classProbabilities, + double& splitInfo, AuxiliarySplitInfo& aux); /** diff --git a/src/mlpack/methods/decision_tree/all_categorical_split_impl.hpp b/src/mlpack/methods/decision_tree/all_categorical_split_impl.hpp index 00135625ab..87da7b3d22 100644 --- a/src/mlpack/methods/decision_tree/all_categorical_split_impl.hpp +++ b/src/mlpack/methods/decision_tree/all_categorical_split_impl.hpp @@ -16,17 +16,18 @@ namespace mlpack { namespace tree { template -template +template double AllCategoricalSplit::SplitIfBetter( const double bestGain, const VecType& data, const size_t numCategories, - const arma::Row& labels, + const arma::Row& labels, + const size_t begin, const size_t numClasses, const WeightVecType& weights, const size_t minimumLeafSize, const double minimumGainSplit, - arma::vec& classProbabilities, + double& splitInfo, AuxiliarySplitInfo& /* aux */) { // Count the number of elements in each potential child. @@ -58,7 +59,7 @@ double AllCategoricalSplit::SplitIfBetter( // Calculate the gain of the split. First we have to calculate the labels // that would be assigned to each child. arma::uvec childPositions(numCategories, arma::fill::zeros); - std::vector> childLabels(numCategories); + std::vector> childLabels(numCategories); std::vector> childWeights(numCategories); for (size_t i = 0; i < numCategories; ++i) { @@ -75,12 +76,12 @@ double AllCategoricalSplit::SplitIfBetter( if (UseWeights) { - childLabels[category][childPositions[category]] = labels[i]; + childLabels[category][childPositions[category]] = labels[begin + i]; childWeights[category][childPositions[category]++] = weights[i]; } else { - childLabels[category][childPositions[category]++] = labels[i]; + childLabels[category][childPositions[category]++] = labels[begin + i]; } } @@ -99,9 +100,8 @@ double AllCategoricalSplit::SplitIfBetter( if (overallGain > bestGain + minimumGainSplit + epsilon) { - // This is better, so set up the class probabilities vector and return. - classProbabilities.set_size(1); - classProbabilities[0] = numCategories; + // This is better, so store it in splitInfo and return. + splitInfo = numCategories; return overallGain; } diff --git a/src/mlpack/methods/decision_tree/decision_tree_impl.hpp b/src/mlpack/methods/decision_tree/decision_tree_impl.hpp index e4cd77851b..8f04e9f57c 100644 --- a/src/mlpack/methods/decision_tree/decision_tree_impl.hpp +++ b/src/mlpack/methods/decision_tree/decision_tree_impl.hpp @@ -642,15 +642,17 @@ double DecisionTree(bestGain, data.cols(begin, begin + count - 1).row(i), datasetInfo.NumMappings(i), - labels.subvec(begin, begin + count - 1), + labels, + begin, numClasses, UseWeights ? weights.subvec(begin, begin + count - 1) : weights, minimumLeafSize, minimumGainSplit, - classProbabilities, + classProbabilities[0], *this); } else if (datasetInfo.Type(i) == data::Datatype::numeric) diff --git a/src/mlpack/tests/decision_tree_test.cpp b/src/mlpack/tests/decision_tree_test.cpp index 67c926ef04..45eb350523 100644 --- a/src/mlpack/tests/decision_tree_test.cpp +++ b/src/mlpack/tests/decision_tree_test.cpp @@ -583,17 +583,17 @@ TEST_CASE("AllCategoricalSplitSimpleSplitTest", "[DecisionTreeTest]") arma::rowvec weights(labels.n_elem); weights.ones(); - arma::vec classProbabilities; + arma::vec classProbabilities(1); AllCategoricalSplit::AuxiliarySplitInfo aux; // Call the method to do the splitting. const double bestGain = GiniGain::Evaluate(labels, 3, weights); const double gain = AllCategoricalSplit::SplitIfBetter( - bestGain, values, 4, labels, 3, weights, 3, 1e-7, classProbabilities, + bestGain, values, 4, labels, 0, 3, weights, 3, 1e-7, classProbabilities[0], aux); const double weightedGain = AllCategoricalSplit::SplitIfBetter(bestGain, values, 4, - labels, 3, weights, 3, 1e-7, classProbabilities, aux); + labels, 0, 3, weights, 3, 1e-7, classProbabilities[0], aux); // Make sure that a split was made. REQUIRE(gain > bestGain); @@ -619,18 +619,17 @@ TEST_CASE("AllCategoricalSplitMinSamplesTest", "[DecisionTreeTest]") arma::rowvec weights(labels.n_elem); weights.ones(); - arma::vec classProbabilities; + arma::vec classProbabilities(1); AllCategoricalSplit::AuxiliarySplitInfo aux; // Call the method to do the splitting. const double bestGain = GiniGain::Evaluate(labels, 3, weights); const double gain = AllCategoricalSplit::SplitIfBetter( - bestGain, values, 4, labels, 3, weights, 4, 1e-7, classProbabilities, - aux); + bestGain, values, 4, labels, 0, 3, weights, 4, 1e-7, + classProbabilities[0], aux); // Make sure it's not split. REQUIRE(gain == DBL_MAX); - REQUIRE(classProbabilities.n_elem == 0); } /** @@ -652,22 +651,21 @@ TEST_CASE("AllCategoricalSplitNoGainTest", "[DecisionTreeTest]") labels[i + 2] = 2; } - arma::vec classProbabilities; + arma::vec classProbabilities(1); AllCategoricalSplit::AuxiliarySplitInfo aux; // Call the method to do the splitting. const double bestGain = GiniGain::Evaluate(labels, 3, weights); const double gain = AllCategoricalSplit::SplitIfBetter( - bestGain, values, 10, labels, 3, weights, 10, 1e-7, - classProbabilities, aux); + bestGain, values, 10, labels, 0, 3, weights, 10, 1e-7, + classProbabilities[0], aux); const double weightedGain = AllCategoricalSplit::SplitIfBetter(bestGain, values, 10, - labels, 3, weights, 10, 1e-7, classProbabilities, aux); + labels, 0, 3, weights, 10, 1e-7, classProbabilities[0], aux); // Make sure that there was no split. REQUIRE(gain == DBL_MAX); REQUIRE(gain == weightedGain); - REQUIRE(classProbabilities.n_elem == 0); } /**