多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

深度聚类代码库盘点:从DeepCluster到SwAV的实战指南

深度聚类代码库盘点:从DeepCluster到SwAV的实战指南 深度聚类这块网上开源代码确实不少但真正能拿来就跑、跑完还能复现出论文指标的库其实就那几个。很多朋友一开始都是对着论文去搜代码结果不是老版本跑不起来就是PyTorch和TensorFlow版本冲突折腾几天全耗在环境上了。这篇文章我把这些年整理过的深度聚类代码库好好梳理一遍从经典方法到统一框架连带着跑实验时踩过的坑、调参的经验、评估指标的坑一起说清楚。不管是刚入门想找个库练手还是做研究需要横向对比应该都能用得上。1. 为什么需要一份深度聚类代码库清单1.1 深度聚类解决的是哪类刚需问题深度聚类的核心诉求很简单在特征提取和聚类之间来回交替优化让网络自己学习出适合聚类的表征。传统聚类算法像K-means、谱聚类输入的都是人工特征特征质量直接决定聚类上限。深度聚类把特征学习和聚类目标绑到一起通过神经网络在高维空间里直接学特征再用聚类结果反向监督特征学习形成一个闭环。这个闭环带来的好处很实际。图像领域里没有标签的数据占绝大多数人工标注成本高、周期长深度聚类可以在完全不依赖标注的情况下把数据按语义自动分组。比如做商品分类、票据归档、人脸聚类甚至是异常检测里的离群样本发现都能用上这套思路。文本、语音、推荐场景同样适用核心都是先学一个能用的特征再让特征和聚类互相促进。1.2 一份能照着跑的代码库清单为什么重要深度聚类论文的代码质量参差不齐有的官方实现写得非常潦草有的依赖老旧库导致今天装不上有的只在特定的分布式环境下能跑。如果没有一份经过验证的清单你可能花在修bug上的时间比跑实验本身还多。更重要的是不同代码库背后对应了不同的技术路线有的走对比学习路线有的走图引导路线有的走生成式路线有的走统一聚类框架。路线不同适用的数据集规模和显存需求也完全不同。比如在CIFAR-10上用DeepCluster非常流畅但换到ImageNet级别显存和训练时间就是另一个量级。提前了解每个库的适用边界能帮你少走很多弯路。2. 主流深度聚类代码库横向盘点2.1 经典三件套DeepCluster、SwAV与SeLaDeepCluster是很多人接触深度聚类入门的第一站核心思路是对特征做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调到50NMI指标能从0.56掉到0.51非常直接。后续的改进版本像IDFD、MSTSC等也都基于近邻关系做文章代码风格都继承自SCAN理解了SCAN再看这些都会很快。2.3 生成式聚类VaDE与ClusterGANVaDE把变分自编码器和高斯混合模型结合到一起把编码空间里的隐变量建模成混合高斯分布从而实现聚类。它的训练目标函数包含重构损失和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.xPython 3.8或3.9这个组合经过大量实测是兼容性最好的。faiss这块很多人第一次装会踩坑装CPU版本平时用没问题但DeepCluster的K-means在GPU上跑和CPU上跑完全两个速度建议直接从conda安装gpu版本。数据预处理上深度聚类对于增广策略的依赖非常高。对比学习路线的代码库普遍使用SimCLR风格的增广组合随机裁剪、颜色抖动、灰度化、高斯模糊。增广强度不能设得太大否则语义信息被破坏聚类结果直接崩。我自己的经验是CIFAR数据集上颜色抖动强度调到0.4到0.5ImageNet级别调到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是起步最好用512ImageNet规模的数据集用多卡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_workersGPI利用率会长期很低。多开几个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都是错的返工成本极高。最后分享一个我自己一直在用的习惯每跑通一个代码库我都会把它的核心配置文件单独存一份注释版本把每个超参数为什么这么设、改大会有什么影响、改小会有什么影响都写在注释里。下次再回头看的时候不用重新翻论文看注释就全想起来了。深度聚类本身就涉及特征、分布、分配、增广这些耦合因素参数之间互相影响单靠记忆很容易混淆。把调参经验沉淀下来比多跑几次实验的收益更大。
返回列表