diff --git a/src/mlpack/methods/random_forest/random_forest_impl.hpp b/src/mlpack/methods/random_forest/random_forest_impl.hpp index 0277f71ece..106072c8f2 100644 --- a/src/mlpack/methods/random_forest/random_forest_impl.hpp +++ b/src/mlpack/methods/random_forest/random_forest_impl.hpp @@ -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"); }