Refactored AllCategoricalSplit

This commit is contained in:
Rishabh Garg
2021-07-12 10:10:01 +05:30
parent a2c9cd0c2d
commit 37559821bf
4 changed files with 28 additions and 28 deletions
@@ -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<bool UseWeights, typename VecType, typename WeightVecType>
template<bool UseWeights, typename VecType, typename ElemType, typename WeightVecType>
static double SplitIfBetter(
const double bestGain,
const VecType& data,
const size_t numCategories,
const arma::Row<size_t>& labels,
const arma::Row<ElemType>& 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);
/**
@@ -16,17 +16,18 @@ namespace mlpack {
namespace tree {
template<typename FitnessFunction>
template<bool UseWeights, typename VecType, typename WeightVecType>
template<bool UseWeights, typename VecType, typename ElemType, typename WeightVecType>
double AllCategoricalSplit<FitnessFunction>::SplitIfBetter(
const double bestGain,
const VecType& data,
const size_t numCategories,
const arma::Row<size_t>& labels,
const arma::Row<ElemType>& 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<FitnessFunction>::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<arma::Row<size_t>> childLabels(numCategories);
std::vector<arma::Row<ElemType>> childLabels(numCategories);
std::vector<arma::Row<double>> childWeights(numCategories);
for (size_t i = 0; i < numCategories; ++i)
{
@@ -75,12 +76,12 @@ double AllCategoricalSplit<FitnessFunction>::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<FitnessFunction>::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;
}
@@ -642,15 +642,17 @@ double DecisionTree<FitnessFunction,
double dimGain = -DBL_MAX;
if (datasetInfo.Type(i) == data::Datatype::categorical)
{
classProbabilities.set_size(1);
dimGain = CategoricalSplit::template SplitIfBetter<UseWeights>(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)
+10 -12
View File
@@ -583,17 +583,17 @@ TEST_CASE("AllCategoricalSplitSimpleSplitTest", "[DecisionTreeTest]")
arma::rowvec weights(labels.n_elem);
weights.ones();
arma::vec classProbabilities;
arma::vec classProbabilities(1);
AllCategoricalSplit<GiniGain>::AuxiliarySplitInfo aux;
// Call the method to do the splitting.
const double bestGain = GiniGain::Evaluate<false>(labels, 3, weights);
const double gain = AllCategoricalSplit<GiniGain>::SplitIfBetter<false>(
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<GiniGain>::SplitIfBetter<true>(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<GiniGain>::AuxiliarySplitInfo aux;
// Call the method to do the splitting.
const double bestGain = GiniGain::Evaluate<false>(labels, 3, weights);
const double gain = AllCategoricalSplit<GiniGain>::SplitIfBetter<false>(
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<GiniGain>::AuxiliarySplitInfo aux;
// Call the method to do the splitting.
const double bestGain = GiniGain::Evaluate<false>(labels, 3, weights);
const double gain = AllCategoricalSplit<GiniGain>::SplitIfBetter<false>(
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<GiniGain>::SplitIfBetter<true>(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);
}
/**