Refactored AllCategoricalSplit
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
Reference in New Issue
Block a user