PyTorch实战:从零构建DEC算法实现深度聚类全流程

深度嵌入式聚类(Deep Embedded Clustering,DEC)作为无监督学习领域的重要突破,通过结合深度神经网络的特征提取能力和传统聚类算法的优势,在图像、文本等多种数据类型上展现出卓越性能。本文将带您用PyTorch完整实现DEC算法,从自编码器构建到聚类结果可视化,涵盖每个技术细节和实战技巧。

1. 环境准备与核心原理

在开始编码前,我们需要明确DEC算法的两大核心组件:自编码器(Autoencoder)和聚类优化模块。自编码器负责学习数据的低维表示,而聚类模块则在这些表示上执行聚类任务,二者通过KL散度损失实现端到端的联合训练。

基础环境配置

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.optim import Adam
from torch.nn.parameter import Parameter
import numpy as np
from sklearn.cluster import KMeans
import matplotlib.pyplot as plt
from sklearn.manifold import TSNE

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

DEC的核心创新在于其双重优化机制

  1. 通过自编码器的重构损失学习有意义的低维表示
  2. 使用聚类目标函数(KL散度)优化特征空间中的分布

2. 自编码器实现与预训练

自编码器作为DEC的特征提取骨干,其结构设计直接影响聚类效果。我们采用三层编码器-解码器架构:

class AE(nn.Module):
    def __init__(self, n_enc_1=500, n_enc_2=500, n_enc_3=2000, 
                 n_dec_1=2000, n_dec_2=500, n_dec_3=500, 
                 n_input=784, n_z=10):
        super(AE, self).__init__()
        # 编码器部分
        self.enc_1 = nn.Linear(n_input, n_enc_1)
        self.enc_2 = nn.Linear(n_enc_1, n_enc_2)
        self.enc_3 = nn.Linear(n_enc_2, n_enc_3)
        self.z_layer = nn.Linear(n_enc_3, n_z)
        
        # 解码器部分
        self.dec_1 = nn.Linear(n_z, n_dec_1)
        self.dec_2 = nn.Linear(n_dec_1, n_dec_2)
        self.dec_3 = nn.Linear(n_dec_2, n_dec_3)
        self.x_bar_layer = nn.Linear(n_dec_3, n_input)

    def forward(self, x):
        # 编码过程
        enc_h1 = F.relu(self.enc_1(x))
        enc_h2 = F.relu(self.enc_2(enc_h1))
        enc_h3 = F.relu(self.enc_3(enc_h2))
        z = self.z_layer(enc_h3)
        
        # 解码过程
        dec_h1 = F.relu(self.dec_1(z))
        dec_h2 = F.relu(self.dec_2(dec_h1))
        dec_h3 = F.relu(self.dec_3(dec_h2))
        x_bar = self.x_bar_layer(dec_h3)
        
        return x_bar, z

预训练技巧

  • 使用MSE损失函数优化重构能力
  • 学习率设置为1e-3,batch size建议128-256
  • 早停策略防止过拟合
def pretrain_ae(model, dataloader, epochs=100):
    optimizer = Adam(model.parameters(), lr=1e-3)
    for epoch in range(epochs):
        total_loss = 0
        for batch in dataloader:
            optimizer.zero_grad()
            x = batch.to(device)
            x_bar, _ = model(x)
            loss = F.mse_loss(x_bar, x)
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        print(f"Epoch {epoch}, Loss: {total_loss/len(dataloader):.4f}")

3. DEC模型构建与聚类层设计

DEC模型在自编码器基础上引入可训练的聚类中心,通过t分布计算样本与聚类中心的相似度:

class DEC(nn.Module):
    def __init__(self, ae, n_clusters=10, v=1.0):
        super(DEC, self).__init__()
        self.ae = ae
        self.v = v  # t分布自由度参数
        self.cluster_layer = Parameter(torch.Tensor(n_clusters, ae.z_layer.out_features))
        torch.nn.init.xavier_normal_(self.cluster_layer.data)
    
    def forward(self, x):
        x_bar, z = self.ae(x)
        
        # 计算样本与聚类中心的相似度(软分配)
        q = 1.0 / (1.0 + torch.sum(torch.pow(z.unsqueeze(1) - self.cluster_layer, 2), 2) / self.v)
        q = q.pow((self.v + 1.0) / 2.0)
        q = (q.t() / torch.sum(q, 1)).t()  # 归一化为概率分布
        
        return x_bar, q, z

关键参数解析

参数 作用 推荐值
v 控制聚类边界硬度 1.0-3.0
n_clusters 聚类数量 根据数据特性确定
learning rate 训练速率 1e-4到1e-3

4. 训练流程与目标分布优化

DEC的训练分为两个阶段:初始化聚类中心和联合优化。我们使用K-means初始化聚类中心,然后交替优化重构损失和KL散度:

def train_dec(model, dataloader, n_clusters, epochs=200):
    # 初始化聚类中心
    model.eval()
    z_all = []
    with torch.no_grad():
        for batch in dataloader:
            _, z = model.ae(batch.to(device))
            z_all.append(z.cpu())
    z_all = torch.cat(z_all, dim=0)
    kmeans = KMeans(n_clusters=n_clusters, n_init=20)
    y_pred = kmeans.fit_predict(z_all.numpy())
    model.cluster_layer.data = torch.tensor(kmeans.cluster_centers_).to(device)
    
    # 联合训练
    optimizer = Adam(model.parameters(), lr=1e-4)
    for epoch in range(epochs):
        model.train()
        total_loss = 0
        for batch in dataloader:
            x = batch.to(device)
            x_bar, q, _ = model(x)
            
            # 计算目标分布p
            p = target_distribution(q.detach())
            
            # 计算KL散度和重构损失
            kl_loss = F.kl_div(q.log(), p, reduction='batchmean')
            re_loss = F.mse_loss(x_bar, x)
            loss = kl_loss + re_loss
            
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        
        # 每10个epoch评估一次
        if epoch % 10 == 0:
            evaluate_clustering(model, dataloader)

目标分布计算函数

def target_distribution(q):
    weight = q**2 / q.sum(0)
    return (weight.t() / weight.sum(1)).t()

5. 结果可视化与性能评估

可视化是验证聚类效果的重要手段,我们使用t-SNE将高维特征降维到2D空间:

def visualize(z, y_pred, y_true=None):
    tsne = TSNE(n_components=2)
    z_2d = tsne.fit_transform(z.cpu().numpy())
    
    plt.figure(figsize=(12, 6))
    plt.subplot(121)
    plt.scatter(z_2d[:, 0], z_2d[:, 1], c=y_pred)
    plt.title("Predicted Clusters")
    
    if y_true is not None:
        plt.subplot(122)
        plt.scatter(z_2d[:, 0], z_2d[:, 1], c=y_true)
        plt.title("True Labels")
    plt.show()

评估指标实现

from sklearn.metrics import normalized_mutual_info_score, adjusted_rand_score

def evaluate_clustering(model, dataloader, y_true=None):
    model.eval()
    z_all, q_all = [], []
    with torch.no_grad():
        for batch in dataloader:
            _, q, z = model(batch.to(device))
            z_all.append(z.cpu())
            q_all.append(q.cpu())
    
    z_all = torch.cat(z_all, dim=0)
    q_all = torch.cat(q_all, dim=0)
    y_pred = q_all.argmax(1).numpy()
    
    if y_true is not None:
        nmi = normalized_mutual_info_score(y_true, y_pred)
        ari = adjusted_rand_score(y_true, y_pred)
        print(f"NMI: {nmi:.4f}, ARI: {ari:.4f}")
    
    return z_all, y_pred

6. 实战技巧与常见问题解决

在实际实现DEC时,有几个关键点需要特别注意:

1. 数据预处理标准化

  • 图像数据归一化到[0,1]
  • 文本数据使用TF-IDF或词向量
  • 数值特征建议标准化(均值0,方差1)

2. 参数调优指南

问题现象 可能原因 解决方案
聚类结果随机 初始化不稳定 增加K-means的n_init参数
所有样本聚为一类 v值过大 减小v值(尝试0.5-2.0)
训练损失震荡 学习率过高 降低学习率或使用学习率衰减

3. 扩展改进思路

  • 替换自编码器为VAE获得更鲁棒的特征表示
  • 加入对抗训练增强特征判别性
  • 使用图卷积网络处理图结构数据
# 改进版损失函数示例(加入特征紧致性约束)
def improved_loss(x, x_bar, q, z, alpha=0.1):
    kl_loss = F.kl_div(q.log(), target_distribution(q.detach()))
    re_loss = F.mse_loss(x_bar, x)
    # 特征紧致性约束
    compact_loss = torch.mean(torch.norm(z, dim=1))
    return kl_loss + re_loss + alpha * compact_loss

通过本教程,您已经掌握了DEC算法的完整实现流程。在实际项目中,建议先在MNIST或Reuters等标准数据集上验证模型效果,再迁移到特定领域数据。DEC的强大之处在于它将特征学习和聚类统一到一个框架中,避免了传统方法中特征工程与聚类分离的局限性。

Logo

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

更多推荐