Add midpoint split. This often produces better trees than mean split...

This commit is contained in:
Ryan Curtin
2015-05-07 15:10:51 -04:00
parent 34eb8da76f
commit 07a099f2b6
2 changed files with 378 additions and 0 deletions
@@ -0,0 +1,117 @@
/**
* @file midpoint_split.hpp
* @author Yash Vadalia
* @author Ryan Curtin
*
* Definition of MidpointSplit, a class that splits a binary space partitioning
* tree node into two parts using the midpoint of the values in a certain
* dimension. The dimension to split on is the dimension with maximum variance.
*/
#ifndef __MLPACK_CORE_TREE_BINARY_SPACE_TREE_MEAN_SPLIT_HPP
#define __MLPACK_CORE_TREE_BINARY_SPACE_TREE_MEAN_SPLIT_HPP
#include <mlpack/core.hpp>
namespace mlpack {
namespace tree /** Trees and tree-building procedures. */ {
/**
* A binary space partitioning tree node is split into its left and right child.
* The split is done in the dimension that has the maximum width. The points are
* divided into two parts based on the midpoint in this dimension.
*/
template<typename BoundType, typename MatType = arma::mat>
class MidpointSplit
{
public:
/**
* Split the node according to the mean value in the dimension with maximum
* width.
*
* @param bound The bound used for this node.
* @param data The dataset used by the binary space tree.
* @param begin Index of the starting point in the dataset that belongs to
* this node.
* @param count Number of points in this node.
* @param splitDimension This will be filled with the dimension the node is to
* be split on.
* @param splitCol The index at which the dataset is divided into two parts
* after the rearrangement.
*/
static bool SplitNode(const BoundType& bound,
MatType& data,
const size_t begin,
const size_t count,
size_t& splitCol);
/**
* Split the node according to the mean value in the dimension with maximum
* width and return a list of changed indices.
*
* @param bound The bound used for this node.
* @param data The dataset used by the binary space tree.
* @param begin Index of the starting point in the dataset that belongs to
* this node.
* @param count Number of points in this node.
* @param splitDimension This will be filled with the dimension the node is
* to be split on.
* @param splitCol The index at which the dataset is divided into two parts
* after the rearrangement.
* @param oldFromNew Vector which will be filled with the old positions for
* each new point.
*/
static bool SplitNode(const BoundType& bound,
MatType& data,
const size_t begin,
const size_t count,
size_t& splitCol,
std::vector<size_t>& oldFromNew);
private:
/**
* Reorder the dataset into two parts such that they lie on either side of
* splitCol.
*
* @param data The dataset used by the binary space tree.
* @param begin Index of the starting point in the dataset that belongs to
* this node.
* @param count Number of points in this node.
* @param splitDimension The dimension to split the node on.
* @param splitVal The split in dimension splitDimension is based on this
* value.
*/
static size_t PerformSplit(MatType& data,
const size_t begin,
const size_t count,
const size_t splitDimension,
const double splitVal);
/**
* Reorder the dataset into two parts such that they lie on either side of
* splitCol. Also returns a list of changed indices.
*
* @param data The dataset used by the binary space tree.
* @param begin Index of the starting point in the dataset that belongs to
* this node.
* @param count Number of points in this node.
* @param splitDimension The dimension to split the node on.
* @param splitVal The split in dimension splitDimension is based on this
* value.
* @param oldFromNew Vector which will be filled with the old positions for
* each new point.
*/
static size_t PerformSplit(MatType& data,
const size_t begin,
const size_t count,
const size_t splitDimension,
const double splitVal,
std::vector<size_t>& oldFromNew);
};
} // namespace tree
} // namespace mlpack
// Include implementation.
#include "midpoint_split_impl.hpp"
#endif
@@ -0,0 +1,261 @@
/**
* @file midpoint_split_impl.hpp
* @author Yash Vadalia
* @author Ryan Curtin
*
* Implementation of class (MidpointSplit) to split a binary space partition
* tree.
*/
#ifndef __MLPACK_CORE_TREE_BINARY_SPACE_TREE_MIDPOINT_SPLIT_IMPL_HPP
#define __MLPACK_CORE_TREE_BINARY_SPACE_TREE_MIDPOINT_SPLIT_IMPL_HPP
#include "midpoint_split.hpp"
namespace mlpack {
namespace tree {
template<typename BoundType, typename MatType>
bool MidpointSplit<BoundType, MatType>::SplitNode(const BoundType& bound,
MatType& data,
const size_t begin,
const size_t count,
size_t& splitCol)
{
size_t splitDimension = data.n_rows; // Indicate invalid.
double maxWidth = -1;
// Find the split dimension. If the bound is tight, we only need to consult
// the bound's width.
if (bound::BoundTraits<BoundType>::HasTightBounds)
{
for (size_t d = 0; d < data.n_rows; d++)
{
const double width = bound[d].Width();
if (width > maxWidth)
{
maxWidth = width;
splitDimension = d;
}
}
}
else
{
// We must individually calculate bounding boxes.
math::Range* ranges = new math::Range[data.n_rows];
for (size_t i = begin; i < begin + count; ++i)
{
// Expand each dimension as necessary.
for (size_t d = 0; d < data.n_rows; ++d)
{
const double val = data(d, i);
if (val < ranges[d].Lo())
ranges[d].Lo() = val;
if (val > ranges[d].Hi())
ranges[d].Hi() = val;
}
}
// Now, which is the widest?
for (size_t d = 0; d < data.n_rows; d++)
{
const double width = ranges[d].Width();
if (width > maxWidth)
{
maxWidth = width;
splitDimension = d;
}
}
delete[] ranges;
}
if (maxWidth == 0) // All these points are the same. We can't split.
return false;
// Split in the midpoint of that dimension.
double splitVal = bound[splitDimension].Mid();
// Perform the actual splitting. This will order the dataset such that points
// with value in dimension splitDimension less than or equal to splitVal are
// on the left of splitCol, and points with value in dimension splitDimension
// greater than splitVal are on the right side of splitCol.
splitCol = PerformSplit(data, begin, count, splitDimension, splitVal);
return true;
}
template<typename BoundType, typename MatType>
bool MidpointSplit<BoundType, MatType>::SplitNode(const BoundType& bound,
MatType& data,
const size_t begin,
const size_t count,
size_t& splitCol,
std::vector<size_t>& oldFromNew)
{
size_t splitDimension = data.n_rows; // Indicate invalid.
double maxWidth = -1;
// Find the split dimension. If the bound is tight, we only need to consult
// the bound's width.
if (bound::BoundTraits<BoundType>::HasTightBounds)
{
for (size_t d = 0; d < data.n_rows; d++)
{
const double width = bound[d].Width();
if (width > maxWidth)
{
maxWidth = width;
splitDimension = d;
}
}
}
else
{
// We must individually calculate bounding boxes.
math::Range* ranges = new math::Range[data.n_rows];
for (size_t i = begin; i < begin + count; ++i)
{
// Expand each dimension as necessary.
for (size_t d = 0; d < data.n_rows; ++d)
{
const double val = data(d, i);
if (val < ranges[d].Lo())
ranges[d].Lo() = val;
if (val > ranges[d].Hi())
ranges[d].Hi() = val;
}
}
// Now, which is the widest?
for (size_t d = 0; d < data.n_rows; d++)
{
const double width = bound[d].Width();
if (width > maxWidth)
{
maxWidth = width;
splitDimension = d;
}
}
}
if (maxWidth == 0) // All these points are the same. We can't split.
return false;
// Split in the midpoint of that dimension.
double splitVal = bound[splitDimension].Mid();
// Perform the actual splitting. This will order the dataset such that points
// with value in dimension splitDimension less than or equal to splitVal are
// on the left of splitCol, and points with value in dimension splitDimension
// greater than splitVal are on the right side of splitCol.
splitCol = PerformSplit(data, begin, count, splitDimension, splitVal,
oldFromNew);
return true;
}
template<typename BoundType, typename MatType>
size_t MidpointSplit<BoundType, MatType>::PerformSplit(
MatType& data,
const size_t begin,
const size_t count,
const size_t splitDimension,
const double splitVal)
{
// This method modifies the input dataset. We loop both from the left and
// right sides of the points contained in this node. The points less than
// splitVal should be on the left side of the matrix, and the points greater
// than splitVal should be on the right side of the matrix.
size_t left = begin;
size_t right = begin + count - 1;
// First half-iteration of the loop is out here because the termination
// condition is in the middle.
while ((data(splitDimension, left) < splitVal) && (left <= right))
left++;
while ((data(splitDimension, right) >= splitVal) && (left <= right) && (right > 0))
right--;
while (left <= right)
{
// Swap columns.
data.swap_cols(left, right);
// See how many points on the left are correct. When they are correct,
// increase the left counter accordingly. When we encounter one that isn't
// correct, stop. We will switch it later.
while ((data(splitDimension, left) < splitVal) && (left <= right))
left++;
// Now see how many points on the right are correct. When they are correct,
// decrease the right counter accordingly. When we encounter one that isn't
// correct, stop. We will switch it with the wrong point we found in the
// previous loop.
while ((data(splitDimension, right) >= splitVal) && (left <= right))
right--;
}
Log::Assert(left == right + 1);
return left;
}
template<typename BoundType, typename MatType>
size_t MidpointSplit<BoundType, MatType>::PerformSplit(
MatType& data,
const size_t begin,
const size_t count,
const size_t splitDimension,
const double splitVal,
std::vector<size_t>& oldFromNew)
{
// This method modifies the input dataset. We loop both from the left and
// right sides of the points contained in this node. The points less than
// splitVal should be on the left side of the matrix, and the points greater
// than splitVal should be on the right side of the matrix.
size_t left = begin;
size_t right = begin + count - 1;
// First half-iteration of the loop is out here because the termination
// condition is in the middle.
while ((data(splitDimension, left) < splitVal) && (left <= right))
left++;
while ((data(splitDimension, right) >= splitVal) && (left <= right) && (right > 0))
right--;
while (left <= right)
{
// Swap columns.
data.swap_cols(left, right);
// Update the indices for what we changed.
size_t t = oldFromNew[left];
oldFromNew[left] = oldFromNew[right];
oldFromNew[right] = t;
// See how many points on the left are correct. When they are correct,
// increase the left counter accordingly. When we encounter one that isn't
// correct, stop. We will switch it later.
while ((data(splitDimension, left) < splitVal) && (left <= right))
left++;
// Now see how many points on the right are correct. When they are correct,
// decrease the right counter accordingly. When we encounter one that isn't
// correct, stop. We will switch it with the wrong point we found in the
// previous loop.
while ((data(splitDimension, right) >= splitVal) && (left <= right))
right--;
}
Log::Assert(left == right + 1);
return left;
}
} // namespace tree
} // namespace mlpack
#endif