深度聚类:SwAV

一、背景:为什么有 SwAV

在传统图像分类里,比如用 ResNet、YOLO、VGG,我们需要:

“图像 + 标签(cat/dog/flower)”

模型才能学习到特征。

但现实问题是:
🔹 很多数据(尤其是工业检测、医学影像)没有人工标签
🔹 我们希望模型能自己学出特征表示,不靠人工标签。

👉 这就是 自监督学习(Self-Supervised Learning)
SwAV 则属于这个范畴。

二、SwAV 有什么用

SwAV 是一种 自监督学习(Self-Supervised Learning)方法,主要用于 视觉特征学习,目标是不依赖人工标注数据就能学到图像的好表示(feature)。

SwAV 的核心目的就是:

  1. 提取图像的高质量特征

    • 它学到的特征向量能很好地表示图像的语义内容。

    • 特征是 通用的 embedding,可以度量图像相似度。

  2. 用于下游任务

    • 聚类(Clustering):比如 K-Means、DBSCAN,把相似的图像聚到一起。特别说明:聚类适合swav模型的 backbone 输出的 avgpool 特征(shape [1,2048,1,1]),保留原始图像的语义信息,即模型的倒数第二层(avgpool)的输出特征然后在经过L2归一化用来聚类;流程:图片 → SwAV avgpool → L2 归一化 → 聚类向量”

    • 分类(Classification):用少量标注数据微调模型即可达到较高精度。

    • 检索(Retrieval):查找相似图像。

    • 其他计算机视觉任务:如目标检测、语义分割,可以把 SwAV 训练好的特征作为 backbone。


💡 总结大白话

SwAV 就像一个“图像特征提取器”,不用标签也能训练好,训练好后,你可以把它拿来做聚类、分类或者找相似图像。

SwAV 并不是直接预测类别,而是学到一套通用的图像特征,之后可以用这些特征做分类、检索或者聚类。

 与 PyTorch 的关系

  • PyTorch 是框架,你可以用它实现 SwAV、YOLO、ResNet 等模型。

  • SwAV 的实现依赖框架(如 PyTorch 或 TensorFlow),但它本身是算法,不是框架。

三,SWAV的实现流程

训练流程大致理解

  1. 数据增强:对图像生成多种视图(crops)。

  2. 特征提取:CNN(如 ResNet)提取每个视图特征。

  3. 原型分配:把特征映射到一组原型上。

  4. 交换预测:让不同视图的特征在原型空间分布一致。

  5. 梯度优化:更新 CNN + 原型参数。

  6. 得到特征:训练结束后,CNN 可以用来提取任意图像的 embedding。

  7. 下游任务:可以用这些 embedding 做聚类、分类或者检索。


💡 一句话总结

YOLO 是目标检测网络,需要标签;SwAV 是自监督特征学习算法,不需要标签,训练出来的是通用特征向量,之后可以做预测或下游任务。

代码实现

swav官方有现成的

一,直接用官方训练好的swav提取图片的特征

from torchvision import transforms
from PIL import Image
import torch

# 从 facebookresearch/swav 仓库加载 resnet50 SwAV 预训练模型
model = torch.hub.load('facebookresearch/swav:main', 'resnet50')
model.eval()

# 图像预处理
transform = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                         std=[0.229, 0.224, 0.225])
])

img = Image.open('your_image.jpg').convert('RGB')
x = transform(img).unsqueeze(0)  # [1, 3, 224, 224]

# 提取特征
with torch.no_grad():
    features = model(x)  # [1, 2048] embedding
print(features.shape)

二,用自己的数据来训练swav

步骤:

  1. 克隆仓库

git clone https://github.com/facebookresearch/swav.git cd swav

  1. 准备数据集

    • SwAV 支持多种数据集,例如 CIFAR-10、ImageNet 等。

    • 数据需要放在 PyTorch 可以读取的目录结构下:


  1. 训练命令示例

python run_distributed_engines.py config=pretrain/swav/swav_8node_resnet.yaml \
    config.DATA.TRAIN_DIR=/path/to/your/data \
    config.OPTIMIZER.MAX_EPOCH=50

这会从头训练 SwAV 模型,生成 checkpoint 保存到默认路径。

  1. 训练完成后

    • 模型 checkpoint 会在指定目录生成,可以用 extract_features.pytorch.hub 加载来提取特征。

用自己的数据训练swav,这里使用训练好的swav提取特征变了点东西

import torch
from swav.models.resnet50 import resnet50  # 需要从官方 repo 拿 model 定义

# 构建模型
model = resnet50()  

# 加载你训练好的 checkpoint
checkpoint = torch.load("/path/to/your/checkpoint.pth")
model.load_state_dict(checkpoint['state_dict'])  # 注意可能需要改 key
model.eval()
# 图像预处理
transform = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                         std=[0.229, 0.224, 0.225])
])

img = Image.open('your_image.jpg').convert('RGB')
x = transform(img).unsqueeze(0)  # [1, 3, 224, 224]

# 提取特征
with torch.no_grad():
    features = model(x)  # [1, 2048] embedding
print(features.shape)

用于下游聚类操作

from torchvision import transforms
from PIL import Image
import torch

# 从 facebookresearch/swav 仓库加载 resnet50 SwAV 预训练模型
model = torch.hub.load('facebookresearch/swav:main', 'resnet50')

backbone = torch.nn.Sequential(*list(model.children())[:-1])
'''model.children()    获取模型最外层子模块
list(model.children()) 转为列表,方便切片
[:-1]  去掉最后一层(通常是分类层)
torch.nn.Sequential(*list(...))    重新封装成一个新的模型,只包含 backbone'''
backbone.eval()#这样提取的特征就变为avgpool层的特征了,注意有些模型最后一层就是avgpool层例如resnet50w4不过它需要在gpu上运行

model.eval()

# 图像预处理
transform = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                         std=[0.229, 0.224, 0.225])
])

img = Image.open('your_image.jpg').convert('RGB')
x = transform(img).unsqueeze(0)  # [1, 3, 224, 224]

# 提取特征

with torch.no_grad():
  features=backbone(x)#features.shape:[1, 2048, 1, 1]
  features=features.view(features.size(0),-1)#view与reshape用法类似features.shape:[1,2048]
  features = features / features.norm(p=2, dim=1, keepdim=True)#L2归一化

异常检测

孤立森林(isolation  Forest)是一种高效的异常检测算法,它和随机森林类似,但每次选择划分属性和划分点(值)时都是随机的,而不是根据信息增益或基尼指数来选择。

 Isolation Forest算法原理

iForest 属于Non-parametric和unsupervised的方法,即不用定义数学模型也不需要有标记的训练。对于如何查找哪些点是否容易被孤立(isolated),iForest使用了一套非常高效的策略。假设我们用一个随机超平面来切割(split)数据空间(data  space),切一次可以生成两个子空间(例如拿刀切蛋糕一分为二)。之后我们再继续用一个随机超平面来切割每个子空间,循环下去,直到每个子空间里面只有一个数据点为止。直观上来讲,我们可以发现那些密度很高的簇是被切分很多次才会停止切割,但是那些密度很低的点很容易很早就停到一个子空间看了。

  iForest 算法得益于随机森林的思想,与随机森林由大量决策树组成一样,iForest森林也由大量的二叉树组成,iForest 中的树叫 isolation tree,简称 iTree,iTree 树和决策树不太一样,其构建过程也比决策树简单,是一个完全随机的过程。

  假设数据集有 N 条数据,构建一颗 ITree时,从 N条数据中均匀抽样(一般是无放回抽样)出 n 个样本出来,作为这棵树的训练样本。在样本中,随机选出一个特征,并在这个特征的所有值范围内(最小值和最大值之间)随机选一个值,对样本进行二叉划分,将样本中小于该值的划分到节点的左边,大于等于该值的划分到节点的右边。由此得到一个分裂条件和左右两边的数据集,然后分别在左右两边的数据集上重复上面的过程,直到数据集只有一条记录或者达到了树的限定高度。

  由于异常数据较小且特征值和正常数据差别很大。因此,构建 iTree的时候,异常数据离根更近,而正常数据离根更远。一颗ITree的结果往往不可信,iForest算法通过多次抽样,构建多颗二叉树。最后整合所有树的结果,并取平均深度作为最终的输出深度,由此计算数据点的异常分支。(平均深度上面为异常点,下面为正常点)

  下图为iForest 构建 iTree 示例,异常数据点(19, 19)通常离根节点很近。

优缺点

      IForest目前是异常点检测最常用的算法之一,它的优点非常突出,他具有线性时间复杂度。因为是随机森林的方法,所以可以用在含有海量数据的数据集上,通常树的数量越多,算法越稳定。由于每棵树都是互相独立生成的,因此可以部署在大规模分布式系统上来加速运算。

  但是IForest也有一些缺点,比如不适用于特别高维的数据。由于每次切数据空间都是随机选取一个维度和该维度的随机一个特征,建完树后仍然有大量的维度没有被使用,导致算法可靠性降低。此推荐降维后使用

  另外IForest仅对即全局稀疏点敏感,不擅长处理局部的相对稀疏点,这样在某些局部的异常点较多的时候检测可能不是很准。

代码实现

from sklearn.ensemble import IsolationForest
isf=IsolationForest(n_estimators=100,max_samples=256,contamination=0.1)#初始化孤立森林对象

常用的一些参数

  • n_estimators:构建多少个itree,(default=100)指定该森林中生成的随机树数量
  • max_samples:采样数,自动是256,(default='auto)

    用来训练随机数的样本数量,即子采样的大小:

      (1)如果设置的是一个int常数,那么就会从总样本 X 拉取 max_samples个样本生成一棵树iTree

      (2)如果设置的是一个float浮点数,那么就会从总样本 X 拉取 max_samples*X.shape[0] 个样本,X.shape[0] 表示总样本个数

contamination:c(n)默认是0.1,float in (0, 0.5),,取值范围为(0, 0.5),表示异常数据占给定的数据集的比例, 在现实中若无法确定,就用“auto”或通过异常分数自行确定值。

from sklearn.ensemble import IsolationForest
isf=IsolationForest(n_estimators=100,max_samples=256,contamination=0.1)#初始化孤立深林对象
isf.fit(data)#训练孤立森林模型; data类型:样本数×特征数
isf.predict(t_data)#预测; t_data类型:样本数×特征数,返回值每个样本属于哪一类,对于异常值返回 -1,对于内部值返回 1

lof异常检测算法

前面学习到isolation Forest孤立森林算法不适合处理局部相对稀疏点,这里引入了lof异常检测算法,适合局部异常检测,无监督学习

基本思想:

LOF 是一种 无监督异常检测算法,主要用来判断数据点在其邻居中是否“异常”。

  • 如果一个点和周围邻居的密度相比 明显稀疏,它可能是异常点。

  • 如果密度和邻居差不多,就是正常点。

核心思想:相对密度异常,不是全局异常,而是“局部异常”。

可达距离出现的原因是 避免邻居本身是孤立点导致密度计算异常,防止邻居是孤立点导致p的 LRD 偏高

代码实现

from sklearn.neighbors import LocalOutlierFactor
lof=LocalOutlierFactor(n_neighbors=2)
'''初始化lof对象;关键参数n_neighbors也就是k:对小数据集,可以选 2~5
对大数据集,可以选 10~50'''
lof.fit_predict(data)
#异常检测,data类型:样本数×特征数,返回值每个样本属于哪一类,1:正常类,-1:异常类

图挖掘

直观理解
CNN →图片局部像素相似性GCN 图节点与邻居关系的相似性

GCN(有参数的图神经网络模型)

核心思想

  • GCN的核心是“邻居特征聚合”
    每个节点的表示不仅依赖自身特征,还会融合其邻居节点的特征,通过多层传播,节点能感知更广的图结构信息。

图结构信息:

图(Graph)由:

  • 节点(Node):表示实体,如人、网页、物体;

  • 边(Edge):表示关系,如好友关系、超链接、交互;

  • 邻接矩阵(Adjacency Matrix,A):存储节点间的连接关系;

  • 节点特征矩阵(Feature Matrix,X):每行是一个节点的特征向量。

代码实现

1:只实现一层传播

from torch_geometric.nn import GCNConv
Gcn=GCNConv(in_channels=16,out_channels=32)#初始化GCN对象
#in_channels:输入特征的维度,out_channels:输出特征的维度
out=Gcn.forward(x=x,edge_index=edge_index)#x:节点特征矩阵[N, F],edge_index:边索引 [2, E]
#前向传播,执行图卷积;含义:每个节点只聚合 一阶邻居 的特征 + 自己的特征(只实现一层)不会自动感知二阶、三阶邻居
#Gcn.forward()与Gcn()效果一样
#out:聚合了邻居节点特征后的新表示[N,out_channels]

2:多层传播

from torch_geometric.nn import GCNConv
import torch.nn.functional as F
import torch
class MultiLayerGCN(torch.nn.Module):
    def __init__(self, in_dim, hidden_dim, out_dim):
        super().__init__()
        self.conv1 = GCNConv(in_dim, hidden_dim)
        self.conv2 = GCNConv(hidden_dim, out_dim)

    def forward(self, x, edge_index):
        x = self.conv1(x, edge_index)  # 第一层 → 聚合一阶邻居
        x = F.relu(x)
        x = self.conv2(x, edge_index)  # 第二层 → 聚合二阶邻居(间接)
        return x

mu=MultiLayerGCN(16,32,32)#初始化MultiLayerGCN
mu.forward(x=x,edge_index=edge_index)#x:节点特征矩阵[N, F],edge_index:边索引 [2, E]
##前向传播,执行图卷积,实现了2层传播,一般2~3层就够了

GAT(有参数的图神经网络模型)

GAT(Graph Attention Network)本质上就是 GCN 的改进版,核心改进在于引入了注意力机制(Attention),让模型自己学习: 哪些邻居节点更重要、影响更大。

GCN/GAT 必须训练才能用,否则参数随机,预测效果几乎随机。不能拿预训练模型直接用

如果最终目的是特征融合训练不需要标签;如果最终目的是节点分类训练需要标签

代码实现(有官方封装好的)

from torch_geometric.nn.models import GAT

model = GAT(
    in_channels=dataset.num_node_features,#输入特征的维度
    hidden_channels=8,#中间层输出特征的维度
    out_channels=...,#最后一层输出的特征维度
    num_layers=2,#进行几层传播
    heads=8,#作用主要是增强模型的表达能力和稳定性
)

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐