Apply @rcurtin comments

Signed-off-by: Omar Shrit <omar@avontech.fr>
This commit is contained in:
Omar Shrit
2023-12-16 22:18:10 +01:00
parent a27ae5074b
commit dc5ccf4e23
4 changed files with 21 additions and 29 deletions
+4 -6
View File
@@ -47,13 +47,14 @@ namespace mlpack {
* with.
*/
template<typename RangeSearchType = RangeSearch<>,
typename PointSelectionPolicy = OrderedPointSelection,
typename MatType = arma::mat>
typename PointSelectionPolicy = OrderedPointSelection>
class DBSCAN
{
public:
//! Easy access to the MatType.
typedef typename RangeSearchType::Mat MatType;
//! Easy access to Element Type of the matrix.
typedef typename MatType::elem_type ElemType;
typedef typename RangeSearchType::Mat::elem_type ElemType;
/**
* Construct the DBSCAN object with the given parameters. The batchMode
* parameter should be set to false in the case where RAM issues will be
@@ -114,9 +115,6 @@ class DBSCAN
//! Maximum distance between two points to be part of same cluster.
ElemType epsilon;
//! Zero, just a variable holder for zero value, can be f16, f32 or double.
ElemType zero = 0.0;
//! Minimum number of points to be in the epsilon-neighborhood (including
//! itself) for the point to be a core-point.
size_t minPoints;
+14 -20
View File
@@ -19,9 +19,8 @@ namespace mlpack {
/**
* Construct the DBSCAN object with the given parameters.
*/
template<typename RangeSearchType, typename PointSelectionPolicy,
typename MatType>
DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::DBSCAN(
template<typename RangeSearchType, typename PointSelectionPolicy>
DBSCAN<RangeSearchType, PointSelectionPolicy>::DBSCAN(
const ElemType epsilon,
const size_t minPoints,
const bool batchMode,
@@ -40,9 +39,8 @@ DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::DBSCAN(
* Performs DBSCAN clustering on the data, returning number of clusters
* and also the centroid of each cluster.
*/
template<typename RangeSearchType, typename PointSelectionPolicy,
typename MatType>
size_t DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::Cluster(
template<typename RangeSearchType, typename PointSelectionPolicy>
size_t DBSCAN<RangeSearchType, PointSelectionPolicy>::Cluster(
const MatType& data,
MatType& centroids)
{
@@ -58,9 +56,8 @@ size_t DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::Cluster(
* Performs DBSCAN clustering on the data, returning number of clusters,
* the centroid of each cluster and also the list of cluster assignments.
*/
template<typename RangeSearchType, typename PointSelectionPolicy,
typename MatType>
size_t DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::Cluster(
template<typename RangeSearchType, typename PointSelectionPolicy>
size_t DBSCAN<RangeSearchType, PointSelectionPolicy>::Cluster(
const MatType& data,
arma::Row<size_t>& assignments,
MatType& centroids)
@@ -94,9 +91,8 @@ size_t DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::Cluster(
* Performs DBSCAN clustering on the data, returning the number of clusters and
* also the list of cluster assignments.
*/
template<typename RangeSearchType, typename PointSelectionPolicy,
typename MatType>
size_t DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::Cluster(
template<typename RangeSearchType, typename PointSelectionPolicy>
size_t DBSCAN<RangeSearchType, PointSelectionPolicy>::Cluster(
const MatType& data,
arma::Row<size_t>& assignments)
{
@@ -146,9 +142,8 @@ size_t DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::Cluster(
* and can save on RAM usage. It may be slower than the batch search with a
* dual-tree algorithm.
*/
template<typename RangeSearchType, typename PointSelectionPolicy,
typename MatType>
void DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::PointwiseCluster(
template<typename RangeSearchType, typename PointSelectionPolicy>
void DBSCAN<RangeSearchType, PointSelectionPolicy>::PointwiseCluster(
const MatType& data,
UnionFind& uf)
{
@@ -182,7 +177,7 @@ void DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::PointwiseCluster(
visited[index] = true;
// Do the range search for only this point.
rangeSearch.Search(data.col(index), RangeType<ElemType>(zero, epsilon),
rangeSearch.Search(data.col(index), RangeType<ElemType>(ElemType(0.0), epsilon),
neighbors,
distances);
@@ -226,9 +221,8 @@ void DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::PointwiseCluster(
* and also the list of cluster assignments. This can perform search in batch,
* naive search).
*/
template<typename RangeSearchType, typename PointSelectionPolicy,
typename MatType>
void DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::BatchCluster(
template<typename RangeSearchType, typename PointSelectionPolicy>
void DBSCAN<RangeSearchType, PointSelectionPolicy>::BatchCluster(
const MatType& data,
UnionFind& uf)
{
@@ -237,7 +231,7 @@ void DBSCAN<RangeSearchType, PointSelectionPolicy, MatType>::BatchCluster(
std::vector<std::vector<ElemType>> distances;
Log::Info << "Performing range search." << std::endl;
rangeSearch.Train(data);
rangeSearch.Search(RangeType<ElemType>(zero, epsilon), neighbors, distances);
rangeSearch.Search(RangeType<ElemType>(ElemType(0.0), epsilon), neighbors, distances);
Log::Info << "Range search complete." << std::endl;
// See the description of the algorithm in `PointwiseCluster()`. The strategy
@@ -46,7 +46,8 @@ class RangeSearch
public:
//! Convenience typedef.
typedef TreeType<MetricType, RangeSearchStat, MatType> Tree;
//! The type of Matrix.
typedef MatType Mat;
//! The type of element held in MatType.
typedef typename MatType::elem_type ElemType;
+1 -2
View File
@@ -230,8 +230,7 @@ TEST_CASE("Float32OutlierSingleModeTest", "[DBSCANTest]")
EuclideanDistance,
arma::Mat<float>,
FloatKDTree>,
OrderedPointSelection,
arma::Mat<float>> d(0.1, 3, false);
OrderedPointSelection> d(0.1, 3, false);
arma::Row<size_t> assignments;
const size_t clusters = d.Cluster(points, assignments);