深度聚类这块,网上开源代码确实不少,但真正能拿来就跑、跑完还能复现出论文指标的库,其实就那几个。很多朋友一开始都是对着论文去搜代码,结果不是老版本跑不起来,就是PyTorch和TensorFlow版本冲突,折腾几天全耗在环境上了。这篇文章我把这些年整理过的深度聚类代码库好好梳理一遍,从经典方法到统一框架,连带着跑实验时踩过的坑、调参的经验、评估指标的坑,一起说清楚。不管是刚入门想找个库练手,还是做研究需要横向对比,应该都能用得上。
1. 为什么需要一份深度聚类代码库清单
1.1 深度聚类解决的是哪类刚需问题
深度聚类的核心诉求很简单:在特征提取和聚类之间来回交替优化,让网络自己学习出适合聚类的表征。传统聚类算法像K-means、谱聚类,输入的都是人工特征,特征质量直接决定聚类上限。深度聚类把特征学习和聚类目标绑到一起,通过神经网络在高维空间里直接学特征,再用聚类结果反向监督特征学习,形成一个闭环。
这个闭环带来的好处很实际。图像领域里,没有标签的数据占绝大多数,人工标注成本高、周期长,深度聚类可以在完全不依赖标注的情况下,把数据按语义自动分组。比如做商品分类、票据归档、人脸聚类,甚至是异常检测里的离群样本发现,都能用上这套思路。文本、语音、推荐场景同样适用,核心都是"先学一个能用的特征,再让特征和聚类互相促进"。
1.2 一份能"照着跑"的代码库清单为什么重要
深度聚类论文的代码质量参差不齐,有的官方实现写得非常潦草,有的依赖老旧库导致今天装不上,有的只在特定的分布式环境下能跑。如果没有一份经过验证的清单,你可能花在修bug上的时间比跑实验本身还多。
更重要的是,不同代码库背后对应了不同的技术路线:有的走对比学习路线,有的走图引导路线,有的走生成式路线,有的走统一聚类框架。路线不同,适用的数据集规模和显存需求也完全不同。比如在CIFAR-10上用DeepCluster非常流畅,但换到ImageNet级别,显存和训练时间就是另一个量级。提前了解每个库的适用边界,能帮你少走很多弯路。
2. 主流深度聚类代码库横向盘点
2.1 经典三件套:DeepCluster、SwAV与SeLa
DeepCluster是很多人接触深度聚类入门的第一站,核心思路是对特征做K-means得到伪标签,然后用伪标签做分类任务来更新网络,循环往复。PyTorch代码在GitHub上非常容易找到,而且有配套的模型权重和配置文件,我在1080Ti上复现CIFAR-10的实验,大概跑十几个小时后ACC能到0.37左右。要注意的是,它的K-means必须跑在GPU上,CPU会慢到怀疑人生,代码里默认用的是faiss,环境里需要装好。
SwAV的全称是Swap Assignments between Views,思路是对同一张图做两种增广,然后让两个分支输出的聚类分配结果互相预测。它比DeepCluster更高效,不用每次都跑K-means,而是维护一个可学习的prototype矩阵。官方源码基于PyTorch,代码抽象程度比较高,看起来会有点费劲,但阅读价值也高。如果你准备做大规模数据集的预训练,SwAV是首选,Imagenet-1K上做自监督预训练的效果很能打。
SeLa的思路则是通过Sinkhorn-Knopp算法直接求解最优分配问题,不需要额外的聚类网络层,收敛非常稳定。它的官方实现代码量很少,阅读门槛最低,我个人认为是了解深度聚类细节最好的入门代码,肉眼看一遍就能理解分配矩阵和损失函数是怎么耦合的。
2.2 图引导聚类:SCAN与后续改进
SCAN的完整名字是Semantic Clustering by Adopting Nearest Neighbors,它把聚类任务拆成两个阶段:先用对比学习预训练特征,再固定特征,用特征空间里的邻居关系和分类器一起做聚类。Pipeline很清晰,预训练和聚类阶段的代码是分开的,想单独替换成其他预训练权重也很方便。
SCAN复现时有一个显著的坑:聚类阶段对超参很敏感,特别是邻居个数K和阈值,K太大容易把不同类别的样本拉进同一个邻居集合,K太小又学不到语义结构。我在CIFAR-100上试过,K从20调到50,NMI指标能从0.56掉到0.51,非常直接。后续的改进版本像IDFD、MSTSC等也都基于近邻关系做文章,代码风格都继承自SCAN,理解了SCAN再看这些都会很快。
2.3 生成式聚类:VaDE与ClusterGAN
VaDE把变分自编码器和高斯混合模型结合到一起,把编码空间里的隐变量建模成混合高斯分布,从而实现聚类。它的训练目标函数包含重构损失和KL散度,代码实现比较干净,跑MNIST这样的低分辨率数据集非常舒服,显存占用小、收敛快,几分钟就能看结果。但换到复杂图像数据集,重构损失容易导致特征不够判别性,效果会明显下滑。
ClusterGAN的思路则完全不同,用对抗生成的方式同时训练生成器和编码器,在潜空间里迫使不同类别分开。它的训练稳定性比VAE难控制,但优点是生出来的样本可以作为聚类结果的直观验证。看代码的时候别只盯着损失函数,它的判别器结构和特殊设计的混合潜变量输入方式才是精髓,理解了这两个点,基本就掌握了代码的全部逻辑。
2.4 统一工具与算法库:哪个适合你
除了按论文复现的代码,还有一类打包好的深度学习工具库,把多种聚类算法统一进一个框架里。比较典型的有slim-TCA,它把深度聚类和TCA迁移成分分析结合在一起,适合处理迁移场景下的聚类任务;还有基于Pytorch实现的各种benchmark库,把DeepCluster、SwAV、SCAN、PCL、DCCM等算法统一起了接口。
对做横向对比研究的同学来说,这类统一框架价值很大。你不用逐个下载每个算法的原始代码,也不用一个个配环境,框架里通常已经把数据集目录、评估脚本、模型保存逻辑都封装好了。缺点是定制化空间相对小,如果你想改损失函数或者加新的memory mechanism,得先理解框架的抽象层级,否则会有点别扭。对只想快速看结果的工程场景,统一框架确实省时间。
3. 跑通一个深度聚类项目的完整实操
3.1 环境准备与数据预处理的细节
我建议先把环境固定下来,PyTorch 1.10以上,CUDA 11.x,Python 3.8或3.9,这个组合经过大量实测是兼容性最好的。faiss这块很多人第一次装会踩坑,装CPU版本平时用没问题,但DeepCluster的K-means在GPU上跑和CPU上跑完全两个速度,建议直接从conda安装gpu版本。
数据预处理上,深度聚类对于增广策略的依赖非常高。对比学习路线的代码库普遍使用SimCLR风格的增广组合:随机裁剪、颜色抖动、灰度化、高斯模糊。增广强度不能设得太大,否则语义信息被破坏,聚类结果直接崩。我自己的经验是,CIFAR数据集上颜色抖动强度调到0.4到0.5,ImageNet级别调到0.8左右,具体数值可以在验证集上小范围试。
3.2 三步跑通一个最小深度聚类项目
以SCAN为例,完整跑通一个实验只需要三个步骤。
第一步,下载代码和数据,把数据目录结构整理成ImageFolder格式,也就是每个类别一个文件夹。虽然聚类本身不需要标签,但验证评估的时候需要真实标签,所以数据目录里要有ground truth文件夹。
第二步,做预训练阶段,训练一个对比学习的特征提取器。这个阶段比较快,CIFAR-10上两三百个epoch基本就够,损失曲线会逐渐下降,但别指望完全收敛后再进入下一步,预训练到差不多就可以停了。
第三步,进入聚类阶段,加载预训练权重,固定特征骨干网络,只训练聚类头。聚类头的设计通常是MLP加上一个小的softmax输出,每个类别对应一个输出神经元。这个阶段需要监控预测类别的熵,如果熵太小说明置信度太高但可能过拟合,熵太大说明聚类还没学起来。
3.3 三类评估指标的读法与计算
深度聚类论文里高频出现三个指标:ACC、NMI、ARI。ACC是无监督聚类准确率,需要把聚类标签和真实标签做最优匹配,通常用匈牙利算法求解,然后计算匹配后的准确率。它反映的是聚类结果在类别层面的正确程度。
NMI是归一化互信息,衡量两个标签分配之间的信息一致性,对聚类的纯度和完备性比较均衡。它的值域在0到1之间,值越大说明聚类结果和真实标签越吻合。ARI是调整兰德指数,会校正掉随机分配带来的偶然一致,所以数值通常看起来比较小,0.4以上的ARI已经是相当好的结果。
这三个指标都有现成的实现,大部分评估代码直接调用sklearn的metrics模块就能算。读实验结果的时候我建议三个指标都看,ACC容易受到类别不均衡影响,NMI对簇的大小比例不敏感,ARI则更严格。如果ACC很高但NMI偏低,通常意味着聚类结果过于碎片化。
3.4 调参记录:学习率、batchsize与聚类批次
我自己跑过不少深度聚类实验,调参上总结出一些规律。
学习率的设置对聚类结果影响非常大,特别是聚类阶段。用SGD优化器的话,学习率从0.01到0.001之间要仔细试,SCAN在CIFAR-10上用0.01配合weight decay 0.0005效果不错,但换到更小的数据集就容易震荡。如果发现聚类损失下降后又反弹,大概率是学习率太大,降到原来的五分之一就会稳定很多。
batchsize的选择直接影响BN统计量和对比学习的负样本数量。在显存允许的范围内,batchsize尽量调大,SwAV这类方法对batchsize非常敏感,小batch下的一致性约束会失效。CIFAR-10上256是起步,最好用512;ImageNet规模的数据集,用多卡1024以上才比较稳。
聚类迭代批次这个参数很多代码库里叫crop iterations,控制的是聚类阶段迭代的次数。这个值不能太大也不能太小,太小聚类头没学充分,太大容易聚类过度集中,导致某些类被吞并。我一般设3到5轮,每轮里面再分多个step,观察每个step的聚类分布变化来决定要不要提前停掉。
4. 常见问题与排查技巧实录
4.1 特征崩塌:深度聚类最常见的翻车点
特征崩塌的表现是:网络学出来的特征全部集中在一个很小的空间区域,聚类结果只有一个大类,ACC直接掉到零点几。这个问题在对比学习和聚类联合训练时尤其容易暴露。
排查思路很简单:把特征做PCA降维,可视化成二维散点图,如果所有点挤成一团,基本可以确定崩塌了。解决办法有几个。一是检查增广策略,确保每张图的两次增广不会严重破坏语义。二是看损失函数是不是没有加入均匀性约束,很多方法会引入entropy regularization或者负熵惩罚来拉开特征分布。三是试试把聚类头的输出神经元数量减少,有时候类别数设置太大,模型找不到足够的区分度,就会选择全部塞进一类里逃避学习。
4.2 显存不足与训练过慢
深度聚类常见的显存瓶颈出现在两个地方:一个是K-means在GPU上跑时,需要把全量特征矩阵放到显存里,CIFAR-10的5万张特征还好,如果换成数据规模上百万的数据集,显存需求直接翻几十倍,一般单卡扛不住。另一个是对比学习需要同时保存所有样本的特征表示,标准做法是维护一个非常大的queue或者memory bank,也会很吃显存。
解法通常是降分辨率、降batchsize、换更轻量的骨干网络。ResNet-18和ResNet-50在聚类任务上的性能差距并不夸张,显存紧张的话先用ResNet-18跑通流程,后续再换大模型。还可以用混合精度训练,现在很多代码库都自带amp接口,开启之后显存能降30%到40%,速度也快不少。
训练过慢的问题,大概率卡在数据加载上。如果用了Online增广且没有开CPU多进程num_workers,GPI利用率会长期很低。多开几个worker,配合缓存加载,通常能解决。别小看这个步骤,我见过不少人FP16都开了,但num_workers设成默认值,训练速度就是上不去。
4.3 复现指标对不上怎么办
复现论文指标对不上,绝大多数不是模型的问题,而是细节不一致。最常见的坑包括:预训练用的数据集划分方式不同、增广参数不完全一致、评估时是否固定随机种子、NMI和ARI的计算是否用了调整后的版本。
我的建议是先把代码库里的默认配置完整跑一遍,不要做任何改动,记录指标。然后逐步改动参数,每次只改一个变量,对比指标变化,才能定位到是哪个环节引入了差异。尤其要注意的是随机种子,深度聚类对随机性的敏感性比普通监督学习高很多,同一个参数换个种子,ACC可能波动两个百分点,最好设置多个种子取平均值来比较。
还有一个隐蔽的坑是特征归一化。聚类前的特征向量要不要做L2归一化,不同的代码库有不同的约定,直接影响了K-means和prototype更新的效果。跑对比实验时,所有方法必须统一特征后处理方式,否则得出的对比结论是不公平的。
4.4 代码库选型一页纸建议
根据我的经验,这里给出比较直观的选型建议。
如果你是完全新手,想先跑通一个完整流程感受一下,首推SeLa的官方代码,代码量小、依赖少、逻辑清晰,跑MNIST或CIFAR-10都很轻松。如果你想做研究、必须和SOTA方法对比,SCAN和SwAV都是稳妥选择,前者结构清晰适合改代码,后者效果好适合刷指标。如果你的数据集规模很大,重点看SwAV和PCL这类对比学习路线,因为它们天然支持大规模特征学习。如果你的场景比较复杂,需要处理迁移、域偏移等,slim-TCA这种结合TCA的库会更省事。
另外提一句折腾代码库时的经验:每个库clone下来后,先跑官方提供的shell脚本或者onescript,确认环境能不能把官方结果复现出来,再动任何代码。这一步是排除环境干扰的最有效手段。很多人一上来就改网络结构,结果跑了几天发现连baseline都是错的,返工成本极高。
最后分享一个我自己一直在用的习惯:每跑通一个代码库,我都会把它的核心配置文件单独存一份注释版本,把每个超参数为什么这么设、改大会有什么影响、改小会有什么影响,都写在注释里。下次再回头看的时候,不用重新翻论文,看注释就全想起来了。深度聚类本身就涉及特征、分布、分配、增广这些耦合因素,参数之间互相影响,单靠记忆很容易混淆。把调参经验沉淀下来,比多跑几次实验的收益更大。