Fixed unnecessary copying of trees

This commit is contained in:
Rishabh Garg
2021-03-26 12:54:13 +05:30
parent ff3a35bdef
commit d90069be5d
@@ -476,21 +476,15 @@ double RandomForest<
DimensionSelectionType& dimensionSelector,
const bool warmStart)
{
size_t oldNumTrees = trees.size();
// Reset the forest if we are not doing a warm-start.
if (!warmStart)
trees.clear();
const size_t oldNumTrees = trees.size();
trees.resize(trees.size() + numTrees);
// Convert avgGain to total gain.
double totalGain = avgGain * oldNumTrees;
if (warmStart)
{
// This will extend the vector with untrained trees.
trees.resize(trees.size() + numTrees);
}
else
{
// This will fill the vector with untrained trees.
trees.resize(numTrees);
}
// Train each tree individually.
#pragma omp parallel for reduction( + : totalGain)
for (omp_size_t i = 0; i < numTrees; ++i)
@@ -503,45 +497,39 @@ double RandomForest<
bootstrapLabels, bootstrapWeights);
Timer::Stop("bootstrap");
// Now build the decision tree.
DecisionTreeType tmpTree;
Timer::Start("train_tree");
if (UseWeights)
{
if (UseDatasetInfo)
{
totalGain += tmpTree.Train(bootstrapDataset, datasetInfo,
bootstrapLabels, numClasses, bootstrapWeights, minimumLeafSize,
minimumGainSplit, maximumDepth, dimensionSelector);
totalGain += trees[oldNumTrees + i].Train(bootstrapDataset,
datasetInfo, bootstrapLabels, numClasses, bootstrapWeights,
minimumLeafSize, minimumGainSplit, maximumDepth,
dimensionSelector);
}
else
{
totalGain += tmpTree.Train(bootstrapDataset, bootstrapLabels, numClasses,
bootstrapWeights, minimumLeafSize, minimumGainSplit, maximumDepth,
dimensionSelector);
totalGain += trees[oldNumTrees + i].Train(bootstrapDataset,
bootstrapLabels, numClasses, bootstrapWeights, minimumLeafSize,
minimumGainSplit, maximumDepth, dimensionSelector);
}
}
else
{
if (UseDatasetInfo)
{
totalGain += tmpTree.Train(bootstrapDataset, datasetInfo,
bootstrapLabels, numClasses, minimumLeafSize, minimumGainSplit,
maximumDepth, dimensionSelector);
totalGain += trees[oldNumTrees + i].Train(bootstrapDataset,
datasetInfo, bootstrapLabels, numClasses, minimumLeafSize,
minimumGainSplit, maximumDepth, dimensionSelector);
}
else
{
totalGain += tmpTree.Train(bootstrapDataset, bootstrapLabels, numClasses,
minimumLeafSize, minimumGainSplit, maximumDepth, dimensionSelector);
totalGain += trees[oldNumTrees + i].Train(bootstrapDataset,
bootstrapLabels, numClasses, minimumLeafSize, minimumGainSplit,
maximumDepth, dimensionSelector);
}
}
// Storing the trained tree at the desired index.
if (warmStart)
trees[oldNumTrees + i] = tmpTree;
else
trees[i] = tmpTree;
Timer::Stop("train_tree");
}