24/11/5 算法笔记 DBSCAN聚类算法

DBSCAN(Density-Based Spatial Clustering of Applications with Noise)是一种基于密度的聚类算法,它能够在具有噪声的空间数据库中发现任意形状的簇。

DBSCAN密度聚类思想

DBSCAN的聚类定义很简单:由密度可达关系导出的最大密度相连的样本集合,即为我们最终聚类的一个类别,或者说一个簇。

这个DBSCAN的簇里面可以有一个或者多个核心对象。如果只有一个核心对象,则簇里其他的非核心对象样本都在这个核心对象的ϵ \epsilonϵ-邻域里;如果有多个核心对象,则簇里的任意一个核心对象的ϵ \epsilonϵ-邻域中一定有一个其他的核心对象,否则这两个核心对象无法密度可达。这些核心对象的ϵ \epsilonϵ-邻域里所有的样本的集合组成的一个D B S C A N DBSCANDBSCAN聚类簇。

那么怎么才能找到这样的簇样本集合呢?DBSCAN使用的方法很简单,它任意选择一个没有类别的核心对象作为种子,然后找到所有这个核心对象能够密度可达的样本集合,即为一个聚类簇。接着继续选择另一个没有类别的核心对象去寻找密度可达的样本集合,这样就得到另一个聚类簇。一直运行到所有核心对象都有类别为止。

下面是一个纯用python实现的DBSCAN算法:

复制代码
import numpy as np

class DBSCAN:
    def __init__(self,eps,min_samples):
        self.eps = eps
        self.min_samples = min_samples

    def fit(self,x):
        self.labels_ = np.full(shape = x.shape[0],fill_value = -1,dtype = int)
#-1表示噪声点
        self.cluster_id_ = 0  #簇标志
        self.x_ = x
        
        for i in range(x.shape[0]):
            if self.labels_[i] != -1:
                continue
            self._expand_cluster(i)
    def _expand_cluster(self,point_idx):
        neighbours = self._region_query(self.x_[point_idx])
        if len(neighbors) < self.min_samples:
            self.labels_[point_idx] = -1  # 标记为噪声点
            return
        
        self.cluster_id_ += 1
        self.labels_[point_idx] = self.cluster_id_
        neighbors_queue = list(neighbors)

        while len(neighbors_queue) > 0:
            neighbor = neighbors_queue.pop(0) #取出邻居点
            if self.labels_[neighbor] == -1:   #检查邻居点的标签
                self.labels_[neighbor] = self.cluster_id_
            if len(self._region_query(self.X_[neighbor])) >= self.min_samples: #区域查询
                for next_neighbor in self._region_query(self.X_[neighbor]): #扩展邻居点
                    if self.labels_[next_neighbor] == -1:
                        self.labels_[next_neighbor] = self.cluster_id_
                        neighbors_queue.append(next_neighbor)


    def _region_query(self,point):#用于查找给定点 point 在指定半径 eps 内的邻居点的核心操作
        return np.where((self.x_ - point) @ (self.x_point)<(self.eps**2)[0]

    
    def get_labels(self):
        return self.labels_

    if __name__ == "__main__":
    np.random.seed(0)
    X = np.random.rand(100, 2)

    dbscan = DBSCAN(eps=0.3, min_samples=5)
    dbscan.fit(X)

    labels = dbscan.get_labels()

    import matplotlib.pyplot as plt

    plt.scatter(X[:, 0], X[:, 1], c=labels, cmap='viridis', marker='o', s=50)
    plt.title('DBSCAN Clustering')
    plt.show()

下面解释每段代码:

fit 方法

复制代码
def fit(self, X):
    self.labels_ = np.full(shape=X.shape[0], fill_value=-1, dtype=int)  # -1 表示噪声点
    self.cluster_id_ = 0
    self.X_ = X

    for i in range(X.shape[0]):
        if self.labels_[i] != -1:
            continue
        self._expand_cluster(i)

fit 方法是 DBSCAN 算法的主要入口点,它初始化聚类标签数组 labels_,所有点最初都被标记为 -1(表示噪声点)。然后,它遍历每个点,如果点还没有被访问过(即标签仍然是 -1),则调用 _expand_cluster 方法来扩展以该点为中心的聚类。

_expand_cluster 方法

复制代码
def _expand_cluster(self, point_idx):
    neighbors = self._region_query(self.X_[point_idx])
    if len(neighbors) < self.min_samples:
        self.labels_[point_idx] = -1  # 标记为噪声点
        return

    self.cluster_id_ += 1
    self.labels_[point_idx] = self.cluster_id_
    neighbors_queue = list(neighbors)

    while len(neighbors_queue) > 0:
        neighbor = neighbors_queue.pop(0)
        if self.labels_[neighbor] == -1:
            self.labels_[neighbor] = self.cluster_id_
        if len(self._region_query(self.X_[neighbor])) >= self.min_samples:
            for next_neighbor in self._region_query(self.X_[neighbor]):
                if self.labels_[next_neighbor] == -1:
                    self.labels_[next_neighbor] = self.cluster_id_
                    neighbors_queue.append(next_neighbor)

_expand_cluster 方法用于扩展聚类。它首先查询给定点的邻居,如果邻居数量少于 min_samples,则将该点标记为噪声点。否则,它将点标记为一个新的聚类,并将其邻居添加到队列中。然后,它迭代地处理邻居的邻居,直到队列为空。

_region_query 方法,(核心)

复制代码
def _region_query(self, point):
    return np.where((self.X_ - point) @ (self.X_ - point) < self.eps**2)[0]

用于找到给定点在 eps 距离内的邻居点。它通过计算点与所有其他点之间的欧氏距离,并返回距离小于 eps 的点的索引。

get_labels 方法

复制代码
def get_labels(self):
    return self.labels_

get_labels 方法返回聚类结果的标签数组

相关推荐
qystca28 分钟前
蓝桥云客---九宫幻方
算法·深度优先·图论
WDeLiang31 分钟前
Flask学习笔记 - 模板渲染
笔记·学习·flask
明月清了个风1 小时前
数据结构与算法学习笔记----贪心区间问题
笔记·学习·算法·贪心算法
努力毕业的小土博^_^1 小时前
【EI/Scopus双检索】2025年4月光电信息、传感云、边缘计算、光学成像、物联网、智慧城市、新材料国际学术盛宴来袭!
人工智能·神经网络·物联网·算法·智慧城市·边缘计算
因为奋斗超太帅啦1 小时前
MySQL学习笔记(一)——MySQL下载安装配置
笔记·学习·mysql
神里流~霜灭1 小时前
数据结构:二叉树(三)·(重点)
c语言·数据结构·c++·算法·二叉树·红黑树·完全二叉树
网安秘谈1 小时前
非对称加密技术深度解析:从数学基础到工程实践
算法
luckyme_1 小时前
leetcode-代码随想录-哈希表-有效的字母异位词
算法·leetcode·散列表
zh_xuan2 小时前
LeeCode 57. 插入区间
c语言·开发语言·数据结构·算法
aoxiang_ywj2 小时前
【Linux】内核驱动学习笔记(二)
linux·笔记·学习