对比学习,很早就听说过这个概念。组里也有人做CLIP相关的工作,看现在的顶会论文,也是热点趋势之一了。近期清华大学发表2026年第一篇science正刊论文:《Deep contrastive learning enables genome-wide virtual screening》。我不清楚虚拟筛选的研究背景,但是看到核心的算法是使用到对比学习。有点引起我兴趣。
在这里插入图片描述
于是这里简单整理一下对比学习的初步实践,后续有时间再整理对比学习的前沿研究。

什么是对比学习?

对比学习(Contrastive Learning)是一类自监督 / 弱监督表征学习方法,核心是通过区分正负样本对,让模型在特征空间中拉近相似样本、推远不相似样本,从而学到数据的内在结构与鲁棒表示,大幅减少对人工标注的依赖。它在计算机视觉、自然语言处理等领域应用广泛,是解决数据稀缺问题的关键技术之一。

  1. 核心逻辑:模型不直接学习标签映射,而是通过成对样本的相似性判断来优化特征表示 —— 对同一实例的不同 “视角”(正样本对)拉近距离,对不同实例(负样本对)推远距离。
  2. 学习目标:在嵌入空间中,最大化正样本对的相似度、最小化负样本对的相似度,使特征具备更强的区分力与泛化能力。
    在这里插入图片描述

Question 1:在对比学习中,什么才算是正样本或者负样本?不一定是做分类任务才用到对比学习吧?
A1:正负样本的定义和任务类型没有必然绑定。

  1. 对比学习不只是用于分类任务,它的核心是通过样本间的相似性关系学习通用表征,适配检索、匹配、聚类等多种任务;
  2. 正负样本的划分,本质是基于 “语义或实例层面的相似性”,而非标签 —— 不同任务场景下,判断相似性的标准完全不同。

Question 2:对比学习最早源自于哪篇论文?研究是如何发展的?
A2:核心对比学习框架的奠基性论文如下:

  1. Sumit Chopra、Raia Hadsell 和 Yann LeCun 发表于 CVPR 2005 的 《Dimensionality Reduction by Learning an Invariant Mapping》,首次系统性提出对比学习的核心框架与对比损失函数(Contrastive Loss),通过学习数据的不变映射,拉近相似样本、拉远不相似样本,成为后续对比学习的核心理论基础。该论文将对比学习用于降维任务,为自监督表示学习提供了重要范式。(论文解读)
  2. Contrastive Predictive Coding》(Aaron van den Oord 等,2018):提出对比预测编码(CPC),通过预测未来样本的上下文表示,实现序列数据的无监督表示学习,为现代对比学习在语音、文本等序列数据领域的应用奠定了基础。(论文解读链接)
  3. Unsupervised Feature Learning via Non-Parametric Instance Discrimination》(Wu 等,2018):提出实例判别任务,引入 Memory Bank 存储负样本表征,解决了大规模数据下负样本存储问题,推动了对比学习在计算机视觉领域的广泛应用。(论文解读链接)

对比学习的初步实践

这里给出三个简单的示例来说明对比学习的研究思想。

1. MINST对比学习特征聚类
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from torchvision.transforms import RandomApply, RandomResizedCrop, RandomHorizontalFlip
import numpy as np
import matplotlib.pyplot as plt
from sklearn.decomposition import PCA

# ===================== 1. 真实图像的数据增强(符合对比学习范式) =====================
class ContrastiveAugmentation:
    """针对MNIST的对比学习数据增强:生成同一样本的两个不同视图"""
    def __init__(self, size=28):
        self.aug = transforms.Compose([
            RandomResizedCrop(size=(size, size), scale=(0.8, 1.0)),  # 随机裁剪+缩放
            RandomHorizontalFlip(p=0.5),  # 随机水平翻转
            RandomApply([transforms.Lambda(lambda x: x + torch.randn_like(x) * 0.05)], p=0.5),  # 随机加噪声
            transforms.Normalize((0.1307,), (0.3081,))  # MNIST标准化
        ])
    
    def __call__(self, x):
        # 对同一张图生成两个不同的增强视图
        x1 = self.aug(x)
        x2 = self.aug(x)
        return x1, x2

# ===================== 2. 基于CNN的编码器(真实视觉任务的编码器) =====================
class MNISTEncoder(nn.Module):
    """轻量CNN编码器,适配MNIST图像"""
    def __init__(self, embedding_dim=64):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(1, 16, kernel_size=3, stride=2, padding=1),  # [1,28,28]→[16,14,14]
            nn.ReLU(),
            nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1), # [16,14,14]→[32,7,7]
            nn.ReLU(),
            nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1), # [32,7,7]→[64,4,4]
            nn.ReLU(),
        )
        self.fc = nn.Linear(64 * 4 * 4, embedding_dim)  # 展平后映射到嵌入维度
    
    def forward(self, x):
        x = self.conv(x)
        x = x.flatten(1)  # [batch, 64,4,4] → [batch, 64*4*4]
        x = self.fc(x)
        return F.normalize(x, dim=1)  # 归一化,方便计算余弦相似度

# ===================== 3. 投影头 + InfoNCE损失(复用核心逻辑) =====================
class ProjectionHead(nn.Module):
    def __init__(self, input_dim=64, hidden_dim=64, output_dim=32):
        super().__init__()
        self.mlp = nn.Sequential(
            nn.Linear(input_dim, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, output_dim)
        )
    
    def forward(self, x):
        return F.normalize(self.mlp(x), dim=1)

def info_nce_loss(features, temperature=0.07):
    """标准InfoNCE损失(适配MNIST批次)"""
    batch_size = features.shape[0] // 2
    # 计算余弦相似度矩阵
    sim_matrix = torch.mm(features, features.t())  # 归一化后,点积=余弦相似度
    
    # 构建正负样本掩码
    pos_mask = torch.zeros((2*batch_size, 2*batch_size), dtype=bool, device=features.device)
    pos_mask[range(batch_size), range(batch_size, 2*batch_size)] = True
    pos_mask[range(batch_size, 2*batch_size), range(batch_size)] = True
    
    neg_mask = ~pos_mask & ~torch.eye(2*batch_size, dtype=bool, device=features.device)
    
    # 计算损失
    pos_sim = sim_matrix[pos_mask].view(2*batch_size, 1)
    neg_sim = sim_matrix[neg_mask].view(2*batch_size, -1)
    
    logits = torch.cat([pos_sim, neg_sim], dim=1) / temperature
    labels = torch.zeros(2*batch_size, dtype=torch.long, device=features.device)
    loss = F.cross_entropy(logits, labels)
    return loss

# ===================== 4. 结果分析:可视化+定量评估 =====================
def visualize_mnist_features(encoder, test_loader, device, num_samples=100):
    """可视化MNIST特征:不同数字用不同颜色,看聚类效果"""
    encoder.eval()
    all_features = []
    all_labels = []
    
    with torch.no_grad():
        for i, (x, y) in enumerate(test_loader):
            if len(all_features) >= num_samples:
                break
            x = x.to(device)
            feat = encoder(x)  # 提取编码器特征
            all_features.append(feat.cpu().numpy())
            all_labels.append(y.numpy())
    
    # 合并并降维
    all_features = np.concatenate(all_features, axis=0)[:num_samples]
    all_labels = np.concatenate(all_labels, axis=0)[:num_samples]
    
    pca = PCA(n_components=2)
    feat_2d = pca.fit_transform(all_features)
    
    # 绘图:不同数字不同颜色
    plt.figure(figsize=(10, 8))
    plt.rcParams['font.sans-serif'] = ['SimHei']  # 指定默认字体
    plt.rcParams['axes.unicode_minus'] = False    # 解决负号显示问题
    colors = plt.cm.tab10(np.linspace(0, 1, 10))
    for digit in range(10):
        mask = all_labels == digit
        plt.scatter(feat_2d[mask, 0], feat_2d[mask, 1], color=colors[digit], label=f"数字{digit}", s=50)
    
    plt.title("MNIST对比学习特征聚类(不同数字颜色不同)", fontsize=14)
    plt.xlabel("PCA维度1")
    plt.ylabel("PCA维度2")
    plt.legend()
    plt.show()

def calculate_similarity_mnist(encoder, aug, test_loader, device, num_batches=5):
    """计算MNIST的正负样本相似度:同数字=正,不同数字=负"""
    encoder.eval()
    pos_sim_list = []
    neg_sim_list = []
    
    with torch.no_grad():
        for i, (x, y) in enumerate(test_loader):
            if i >= num_batches:
                break
            x = x.to(device)
            y = y.to(device)
            
            # 生成两个增强视图并提取特征
            x1, x2 = aug.aug(x), aug.aug(x)  # 同一样本的两个视图(正样本)
            feat1 = encoder(x1)
            feat2 = encoder(x2)
            
            # 1. 正样本相似度:同一样本的两个视图
            pos_sim = F.cosine_similarity(feat1, feat2, dim=1)
            pos_sim_list.extend(pos_sim.cpu().numpy())
            
            # 2. 负样本相似度:不同数字的样本
            # 打乱批次内样本,构造不同数字的负样本对
            idx_shuffle = torch.randperm(len(x))
            y_shuffle = y[idx_shuffle]
            mask = y != y_shuffle  # 只保留不同数字的样本对
            if mask.sum() == 0:
                continue
            neg_feat1 = feat1[mask]
            neg_feat2 = feat2[idx_shuffle][mask]
            neg_sim = F.cosine_similarity(neg_feat1, neg_feat2, dim=1)
            neg_sim_list.extend(neg_sim.cpu().numpy())
    
    avg_pos = np.mean(pos_sim_list)
    avg_neg = np.mean(neg_sim_list)
    return avg_pos, avg_neg

# ===================== 5. 完整训练流程 =====================
def main():
    # 设备配置
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"使用设备:{device}")
    
    # 超参数
    batch_size = 128
    epochs = 100
    lr = 0.001
    embedding_dim = 64
    
    # 1. 加载MNIST真实数据集
    transform_base = transforms.Compose([transforms.ToTensor()])
    train_dataset = datasets.MNIST(
        root="./data", train=True, download=True, transform=transform_base
    )
    test_dataset = datasets.MNIST(
        root="./data", train=False, download=True, transform=transform_base
    )
    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=0)
    test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
    
    # 2. 初始化组件
    aug = ContrastiveAugmentation(size=28)
    encoder = MNISTEncoder(embedding_dim=embedding_dim).to(device)
    projection_head = ProjectionHead(input_dim=embedding_dim).to(device)
    optimizer = torch.optim.Adam(
        list(encoder.parameters()) + list(projection_head.parameters()),
        lr=lr, weight_decay=1e-4
    )
    
    # 3. 训练前评估(初始状态)
    print("\n===== 训练前 =====")
    avg_pos_pre, avg_neg_pre = calculate_similarity_mnist(encoder, aug, test_loader, device)
    print(f"同数字(正样本)平均相似度:{avg_pos_pre:.4f}")
    print(f"不同数字(负样本)平均相似度:{avg_neg_pre:.4f}")
    visualize_mnist_features(encoder, test_loader, device, num_samples=500)
    
    # 4. 对比学习训练
    encoder.train()
    projection_head.train()
    for epoch in range(epochs):
        total_loss = 0.0
        for x, _ in train_loader:  # 自监督学习,不需要标签
            x = x.to(device)
            optimizer.zero_grad()
            
            # 生成两个增强视图
            x1, x2 = aug(x)
            # 编码+投影
            feat1 = encoder(x1)
            feat2 = encoder(x2)
            z1 = projection_head(feat1)
            z2 = projection_head(feat2)
            # 拼接特征计算损失
            z = torch.cat([z1, z2], dim=0)
            loss = info_nce_loss(z)
            
            # 反向传播
            loss.backward()
            optimizer.step()
            
            total_loss += loss.item()
        
        avg_loss = total_loss / len(train_loader)
        print(f"\nEpoch [{epoch+1}/{epochs}], Loss: {avg_loss:.4f}")
    
    # 5. 训练后评估
    print("\n===== 训练后 =====")
    avg_pos_post, avg_neg_post = calculate_similarity_mnist(encoder, aug, test_loader, device)
    print(f"同数字(正样本)平均相似度:{avg_pos_post:.4f}")
    print(f"不同数字(负样本)平均相似度:{avg_neg_post:.4f}")
    
    # 6. 可视化训练后的特征聚类
    visualize_mnist_features(encoder, test_loader, device, num_samples=500)

if __name__ == "__main__":
    main()

实验结果如下:
===== 训练前 =====
同数字(正样本)平均相似度:0.8967
不同数字(负样本)平均相似度:0.8017
===== 训练后 =====
同数字(正样本)平均相似度:0.9690
不同数字(负样本)平均相似度:0.2438
在这里插入图片描述
在这里插入图片描述
上面两个图是训练前后的特征聚类可视化图。

2.文本语义特征学习
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, Dataset
import numpy as np
import random
from sklearn.metrics.pairwise import cosine_similarity
from sklearn.decomposition import PCA
import jieba
import matplotlib.pyplot as plt

plt.rcParams["font.family"] = ["SimHei"]
plt.rcParams['axes.unicode_minus'] = False

def set_seed(seed=42):
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    np.random.seed(seed)
    random.seed(seed)

set_seed(42)

# ===================== 数据 =====================
RAW_TEXTS = {
    "猫科": ["猫是小型哺乳动物,作为宠物饲养", "猫咪喜欢吃鱼和老鼠,属于猫科", "老虎是大型猫科动物,生活在森林", "狮子是群居猫科动物,万兽之王"],
    "犬科": ["狗是人类的好朋友,属于犬科", "犬的嗅觉灵敏,常作为工作犬", "哈士奇是雪橇犬,性格活泼", "金毛犬温顺,适合陪伴老人"],
    "天体": ["太阳是太阳系中心的恒星", "月亮是地球的天然卫星,不发光", "地球是太阳系的行星,有生命", "火星是红色行星,适合探索"],
    "水果": ["苹果是常见水果,富含维生素", "香蕉是热带水果,能补充钾", "橙子酸甜可口,富含维C", "草莓是浆果,味道鲜美"]
}

TEXT_LIST = []
TEXT_LABELS = []
for label_idx, (category, texts) in enumerate(RAW_TEXTS.items()):
    TEXT_LIST.extend(texts)
    TEXT_LABELS.extend([label_idx] * len(texts))

# ✅ 温和增强:只做同义词替换 + 轻微删除
def text_augmentation(text):
    words = list(jieba.cut(text))
    synonym_dict = {
        "猫": ["猫咪"], "狗": ["犬"], "太阳": ["日"], "月亮": ["月球"],
        "苹果": ["红苹果"], "香蕉": ["蕉"], "老虎": ["虎"], "狮子": ["狮"]
    }
    # 同义词替换(30%概率)
    for i in range(len(words)):
        if random.random() < 0.3 and words[i] in synonym_dict:
            words[i] = random.choice(synonym_dict[words[i]])
    # 轻微删除(最多删1个词,10%概率每个词)
    if len(words) > 3:
        words = [w for w in words if random.random() > 0.1]
    return "".join(words)

# ===================== Dataset =====================
class TextDataset(Dataset):
    def __init__(self, texts, labels, vocab, max_len=15):
        self.texts = texts
        self.labels = labels
        self.vocab = vocab
        self.max_len = max_len
    
    def __len__(self):
        return len(self.texts)
    
    def text_to_idx(self, text):
        words = list(jieba.cut(text))[:self.max_len]
        idx = [self.vocab.get(w, 0) for w in words]
        if len(idx) < self.max_len:
            idx += [0] * (self.max_len - len(idx))
        return torch.tensor(idx, dtype=torch.long)
    
    def __getitem__(self, idx):
        text = self.texts[idx]
        label = self.labels[idx]
        text_pos = text_augmentation(text)
        return {
            "anchor": self.text_to_idx(text),
            "pos": self.text_to_idx(text_pos),
            "text": text,
            "label": label
        }

# ===================== 简化模型:Avg Embedding =====================
class SimpleTextEncoder(nn.Module):
    def __init__(self, vocab_size, embed_dim=64):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
    
    def forward(self, x):
        # x: [B, L]
        mask = (x != 0).float().unsqueeze(-1)  # [B, L, 1]
        embed = self.embedding(x)              # [B, L, D]
        # 平均非 PAD 词
        sum_embed = torch.sum(embed * mask, dim=1)      # [B, D]
        count = torch.sum(mask, dim=1).clamp(min=1)     # [B, 1]
        avg_embed = sum_embed / count                   # [B, D]
        return F.normalize(avg_embed, dim=1)

# ===================== 标准 InfoNCE Loss (2N×2N) =====================
def info_nce_loss_simple(anchor_z, pos_z, temperature=0.5):
    # anchor_z, pos_z: [B, D]
    device = anchor_z.device
    batch_size = anchor_z.size(0)
    
    # 拼接所有样本
    z = torch.cat([anchor_z, pos_z], dim=0)  # [2B, D]
    sim_matrix = torch.mm(z, z.t()) / temperature  # [2B, 2B]
    
    # 正样本对:(i, i+B) 和 (i+B, i)
    pos_mask = torch.zeros_like(sim_matrix, dtype=torch.bool)
    pos_mask[range(batch_size), range(batch_size, 2*batch_size)] = True
    pos_mask[range(batch_size, 2*batch_size), range(batch_size)] = True
    
    # 排除对角线(自己和自己)
    eye = torch.eye(2*batch_size, dtype=torch.bool, device=device)
    sim_matrix = sim_matrix.masked_fill(eye, float('-inf'))
    
    # 计算 loss
    exp_sim = torch.exp(sim_matrix)
    log_prob = sim_matrix - torch.log(exp_sim.sum(dim=1, keepdim=True))
    loss = -log_prob[pos_mask].mean()
    return loss

# ===================== 评估函数(不变) =====================
def calculate_avg_similarity(encoder, dataloader, device):
    encoder.eval()
    pos_sim_list, neg_sim_list = [], []
    with torch.no_grad():
        for batch in dataloader:
            anchor = batch["anchor"].to(device)
            pos = batch["pos"].to(device)
            labels = batch["label"].to(device)
            a_feat = encoder(anchor)
            p_feat = encoder(pos)
            pos_sim = F.cosine_similarity(a_feat, p_feat, dim=1)
            pos_sim_list.extend(pos_sim.cpu().numpy())
            B = a_feat.size(0)
            for i in range(B):
                for j in range(B):
                    if i != j and labels[i] != labels[j]:
                        s = F.cosine_similarity(a_feat[i:i+1], a_feat[j:j+1]).item()
                        neg_sim_list.append(s)
    return np.mean(pos_sim_list or [0]), np.mean(neg_sim_list or [0])

def visualize_feature_space(encoder, dataloader, device, title):
    encoder.eval()
    feats, labels = [], []
    with torch.no_grad():
        for batch in dataloader:
            f = encoder(batch["anchor"].to(device)).cpu().numpy()
            feats.append(f)
            labels.extend(batch["label"].numpy())
    feats = np.concatenate(feats, axis=0)
    feat_2d = PCA(n_components=2).fit_transform(feats)
    plt.figure(figsize=(10, 8))
    colors = ["#FF6B6B", "#4ECDC4", "#45B7D1", "#96CEB4"]
    for i, name in enumerate(["猫科", "犬科", "天体", "水果"]):
        mask = np.array(labels) == i
        plt.scatter(feat_2d[mask, 0], feat_2d[mask, 1], c=colors[i], label=name, s=120, alpha=0.8)
    plt.title(title, fontsize=16)
    plt.legend()
    plt.grid(alpha=0.3)
    plt.show()

def test_retrieval(encoder, dataloader, query_text, device, top_k=3):
    encoder.eval()
    all_feats, all_texts = [], []
    with torch.no_grad():
        for batch in dataloader:
            f = encoder(batch["anchor"].to(device)).cpu().numpy()
            all_feats.append(f)
            all_texts.extend(batch["text"])
    all_feats = np.concatenate(all_feats, axis=0)
    vocab = dataloader.dataset.vocab
    q_idx = dataloader.dataset.text_to_idx(query_text).unsqueeze(0).to(device)
    # 使用 .detach() 方法
    q_feat = encoder(q_idx).detach().cpu().numpy()
    scores = cosine_similarity(q_feat, all_feats)[0]
    top_idx = np.argsort(-scores)[:top_k]
    print(f"\n🔍 查询: {query_text}")
    for i, idx in enumerate(top_idx):
        print(f"{i+1}. {all_texts[idx]} | 相似度: {scores[idx]:.4f}")

# ===================== 主流程 =====================
def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"💻 Device: {device}")

    # 构建词表
    all_words = [w for t in TEXT_LIST for w in jieba.cut(t)]
    vocab = {"<PAD>": 0}
    for w in all_words:
        if w not in vocab:
            vocab[w] = len(vocab)
    print(f"📚 Vocab size: {len(vocab)}")

    dataset = TextDataset(TEXT_LIST, TEXT_LABELS, vocab)
    dataloader = DataLoader(dataset, batch_size=8, shuffle=True)

    encoder = SimpleTextEncoder(vocab_size=len(vocab), embed_dim=64).to(device)
    optimizer = torch.optim.AdamW(encoder.parameters(), lr=3e-3, weight_decay=1e-4)

    # 训练前
    print("\n📊 Before training:")
    pos0, neg0 = calculate_avg_similarity(encoder, dataloader, device)
    print(f"Positive: {pos0:.4f}, Negative: {neg0:.4f}")
    visualize_feature_space(encoder, dataloader, device, "Before Training")

    # 训练
    epochs = 80
    pos_trend, neg_trend = [], []
    encoder.train()
    for epoch in range(epochs):
        total_loss = 0
        for batch in dataloader:
            a = batch["anchor"].to(device)
            p = batch["pos"].to(device)
            optimizer.zero_grad()
            a_z = encoder(a)
            p_z = encoder(p)
            loss = info_nce_loss_simple(a_z, p_z, temperature=0.5)
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        pos, neg = calculate_avg_similarity(encoder, dataloader, device)
        pos_trend.append(pos)
        neg_trend.append(neg)
        if (epoch+1) % 20 == 0:
            print(f"Epoch {epoch+1}: Loss={total_loss/len(dataloader):.4f}, Pos={pos:.4f}, Neg={neg:.4f}")

    # 训练后
    print("\n📊 After training:")
    pos1, neg1 = calculate_avg_similarity(encoder, dataloader, device)
    print(f"Positive: {pos1:.4f}, Negative: {neg1:.4f}")
    
    # 绘制趋势
    plt.figure(figsize=(10, 6))
    plt.plot(pos_trend, "r-o", label="Positive")
    plt.plot(neg_trend, "b--s", label="Negative")
    plt.title("Similarity Trend")
    plt.xlabel("Epoch")
    plt.ylabel("Cosine Similarity")
    plt.legend()
    plt.grid(alpha=0.3)
    plt.show()

    visualize_feature_space(encoder, dataloader, device, "After Training")
    test_retrieval(encoder, dataloader, "猫是宠物", device)
    test_retrieval(encoder, dataloader, "苹果是水果", device)

if __name__ == "__main__":
    main()

实验结果如下:
💻 Device: cuda
📚 Vocab size: 80
📊 Before training:
Positive: 0.8877, Negative: 0.2345

Epoch 20: Loss=1.2860, Pos=0.9435, Neg=0.0899
Epoch 40: Loss=1.1461, Pos=0.9117, Neg=0.0149
Epoch 60: Loss=1.0713, Pos=0.9219, Neg=-0.0483
Epoch 80: Loss=1.0857, Pos=0.9286, Neg=-0.0650

📊 After training:
Positive: 0.9577, Negative: -0.0669

在这里插入图片描述
在这里插入图片描述
在这里插入图片描述
🔍 查询: 猫是宠物
1. 猫是小型哺乳动物,作为宠物饲养 | 相似度: 0.6358
2. 哈士奇是雪橇犬,性格活泼 | 相似度: 0.2716
3. 草莓是浆果,味道鲜美 | 相似度: 0.2047

🔍 查询: 苹果是水果
4. 苹果是常见水果,富含维生素 | 相似度: 0.6154
5. 香蕉是热带水果,能补充钾 | 相似度: 0.3778
6. 火星是红色行星,适合探索 | 相似度: 0.1858

3.图对比学习

前几年GNN比较火的时候,图对比学习算法经常发在顶会顶刊,最近一年研究热度稍微冷了一些。但还是有很多人研究。这里给出一个基本示例:

import torch
import torch.nn.functional as F
from torch_geometric.datasets import TUDataset
from torch_geometric.loader import DataLoader
from torch_geometric.nn import GCNConv, global_mean_pool, LayerNorm
import networkx as nx
import random
import matplotlib.pyplot as plt
from sklearn.manifold import TSNE
import numpy as np

# 设置中文字体(可选,避免中文乱码)
plt.rcParams['font.sans-serif'] = ['DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False

# 设置随机种子保证可复现性
def set_seed(seed=42):
    random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed(seed)
        torch.cuda.manual_seed_all(seed)

set_seed(42)

# --------------------------
# 1. 可视化工具函数(修复图例+优化配色)
# --------------------------
def plot_graph(data, title="Graph", ax=None):
    """可视化单个图"""
    if ax is None:
        ax = plt.gca()
    G = nx.Graph()
    edge_index = data.edge_index.cpu().numpy()
    G.add_edges_from(edge_index.T)
    pos = nx.spring_layout(G)
    nx.draw(G, pos, ax=ax, node_color="#4CAF50", node_size=200, 
            edge_color="#9E9E9E", linewidths=0.5, font_size=8)
    ax.set_title(title, fontsize=12)
    ax.axis('off')

def plot_augmentation_effect(original_data, aug_data1, aug_data2):
    """可视化图增强效果"""
    fig, axes = plt.subplots(1, 3, figsize=(15, 4))
    plot_graph(original_data, "Original Graph", axes[0])
    plot_graph(aug_data1, "Augmented View 1 (Drop Node/Edge)", axes[1])
    plot_graph(aug_data2, "Augmented View 2 (Drop Node/Edge)", axes[2])
    plt.tight_layout()
    plt.savefig("graph_augmentation_effect.png", dpi=300, bbox_inches='tight')
    plt.show()

def plot_loss_curve(loss_history):
    """可视化损失曲线(带早停标记)"""
    plt.figure(figsize=(10, 5))
    # 原始损失曲线
    plt.plot(range(1, len(loss_history)+1), loss_history, 
             color="#2196F3", linewidth=1, alpha=0.5, label="Raw Loss")
    # 平滑损失曲线(窗口大小5)
    if len(loss_history) >= 5:
        smooth_loss = np.convolve(loss_history, np.ones(5)/5, mode='valid')
        plt.plot(range(3, len(loss_history)-1), smooth_loss, 
                 color="#F44336", linewidth=2, label="Smoothed Loss")
    
    # 标记早停点(Loss稳定轮数)
    if len(loss_history) >= 50:
        # 找到Loss稳定的轮数(连续10轮波动<0.02)
        stable_epoch = None
        for i in range(50, len(loss_history)-10):
            window = loss_history[i:i+10]
            if max(window) - min(window) < 0.02:
                stable_epoch = i+1
                break
        if stable_epoch:
            plt.axvline(x=stable_epoch, color="#FF9800", linestyle='--', 
                        label=f"Stable at Epoch {stable_epoch}")
    
    plt.xlabel("Epoch", fontsize=12)
    plt.ylabel("Average Total Loss", fontsize=12)
    plt.title("Training Loss Curve (Anti-Overfitting)", fontsize=14)
    plt.grid(True, alpha=0.3)
    plt.legend()
    plt.savefig("loss_curve_anti_overfit.png", dpi=300, bbox_inches='tight')
    plt.show()

def plot_feature_tsne(z1, z2, stage):
    """极简清晰版TSNE:聚焦同一图的聚集效果"""
    all_features = torch.cat([z1, z2], dim=0).cpu().detach().numpy()
    num_graphs = z1.shape[0]
    
    # TSNE参数
    perplexity = min(20, num_graphs-1) if num_graphs > 1 else 1
    tsne = TSNE(n_components=2, random_state=42, perplexity=perplexity, 
                learning_rate=300, max_iter=1500)
    tsne_features = tsne.fit_transform(all_features)
    
    # 绘图设置
    plt.figure(figsize=(12, 9))
    # 使用颜色循环(自动适配任意数量图,避免重复混淆)
    colors = plt.cm.get_cmap('hsv', num_graphs)  # 色调循环,每个图对应唯一色调
    marker_view1 = 'o'  # 视图1:实心圆
    marker_view2 = '^'  # 视图2:实心三角
    size = 80
    alpha = 0.8
    
    # 绘制每个图的两个视图
    for i in range(num_graphs):
        # 视图1(Drop Node)
        x1, y1 = tsne_features[i, 0], tsne_features[i, 1]
        plt.scatter(x1, y1, color=colors(i), marker=marker_view1, 
                    s=size, alpha=alpha, edgecolors='black', linewidth=0.3)
        # 视图2(Drop Edge)
        x2, y2 = tsne_features[i + num_graphs, 0], tsne_features[i + num_graphs, 1]
        plt.scatter(x2, y2, color=colors(i), marker=marker_view2, 
                    s=size, alpha=alpha, edgecolors='black', linewidth=0.3)
        # 同一图的连线
        plt.plot([x1, x2], [y1, y2], color=colors(i), linewidth=1.2, alpha=0.6)
    
    # 仅显示“视图类型”图例(聚焦核心效果)
    plt.legend(
        handles=[
            plt.scatter([], [], color='gray', marker=marker_view1, label='View 1 (Drop Node)'),
            plt.scatter([], [], color='gray', marker=marker_view2, label='View 2 (Drop Edge)')
        ],
        fontsize=11, loc="best"
    )
    
    # 图表样式
    plt.xlabel("TSNE Dimension 1", fontsize=12)
    plt.ylabel("TSNE Dimension 2", fontsize=12)
    plt.title(f"Graph Feature TSNE ({stage} Training)", fontsize=14)
    plt.grid(True, alpha=0.3, linestyle='--')
    plt.tight_layout()
    plt.savefig(f"feature_tsne_{stage}_final.png", dpi=300, bbox_inches='tight')
    plt.show()

# --------------------------
# 2. 图增强函数(适度增强)
# --------------------------
def graph_augmentation(data, drop_node_p=0.1, drop_edge_p=0.1):
    """适度增强,保留80%+节点/边"""
    if isinstance(data, list):
        data = data[0]
    
    aug_data = data.clone()
    device = aug_data.x.device
    
    # 节点丢弃
    if drop_node_p > 0 and aug_data.x is not None and aug_data.x.shape[0] > 0:
        num_nodes = aug_data.x.shape[0]
        node_mask = (torch.rand(num_nodes, device=device) > drop_node_p)
        min_nodes = max(1, int(num_nodes * 0.8))
        if node_mask.sum() < min_nodes:
            node_scores = torch.rand(num_nodes, device=device)
            top_k_idx = torch.topk(node_scores, min_nodes).indices
            node_mask = torch.zeros(num_nodes, dtype=torch.bool, device=device)
            node_mask[top_k_idx] = True
        
        aug_data.x = aug_data.x[node_mask]
        if hasattr(aug_data, 'batch') and aug_data.batch is not None:
            aug_data.batch = aug_data.batch[node_mask]
        
        # 更新边索引
        if aug_data.edge_index is not None and aug_data.edge_index.size(1) > 0:
            edge_index = aug_data.edge_index
            mask = node_mask[edge_index[0]] & node_mask[edge_index[1]]
            aug_data.edge_index = edge_index[:, mask]
            
            node_idx = torch.arange(num_nodes, device=device)[node_mask]
            new_idx = torch.zeros(num_nodes, dtype=torch.long, device=device) - 1
            new_idx[node_idx] = torch.arange(len(node_idx), device=device)
            aug_data.edge_index = new_idx[aug_data.edge_index]
    
    # 边扰动
    if drop_edge_p > 0 and aug_data.edge_index is not None and aug_data.edge_index.size(1) > 0:
        edge_mask = (torch.rand(aug_data.edge_index.size(1), device=device) > drop_edge_p)
        min_edges = max(1, int(aug_data.edge_index.size(1) * 0.7))
        if edge_mask.sum() < min_edges:
            edge_scores = torch.rand(aug_data.edge_index.size(1), device=device)
            top_k_idx = torch.topk(edge_scores, min_edges).indices
            edge_mask = torch.zeros(aug_data.edge_index.size(1), dtype=torch.bool, device=device)
            edge_mask[top_k_idx] = True
        aug_data.edge_index = aug_data.edge_index[:, edge_mask]
    
    return aug_data

# --------------------------
# 3. 增强版GCN编码器
# --------------------------
class EnhancedGCNEncoder(torch.nn.Module):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.norm1 = LayerNorm(hidden_channels)
        self.conv2 = GCNConv(hidden_channels, hidden_channels)
        self.norm2 = LayerNorm(hidden_channels)
        self.conv3 = GCNConv(hidden_channels, out_channels)
        self.norm3 = LayerNorm(out_channels)
    
    def forward(self, x, edge_index, batch):
        x = self.conv1(x, edge_index)
        x = self.norm1(x)
        x = F.relu(x)
        x = F.dropout(x, p=0.1, training=self.training)
        
        x = self.conv2(x, edge_index)
        x = self.norm2(x)
        x = F.relu(x)
        
        x = self.conv3(x, edge_index)
        x = self.norm3(x)
        
        x = global_mean_pool(x, batch)
        x = F.normalize(x, p=2, dim=-1)
        return x

# --------------------------
# 4. 带多样性约束的损失函数
# --------------------------
def infonce_loss(z1, z2, temperature=0.4):
    """调整温度系数,增强区分度"""
    batch_size = z1.size(0)
    sim_matrix = torch.mm(z1, z2.t()) / temperature
    
    # 正样本(对角线)
    pos_sim = torch.diag(sim_matrix)
    # 负样本(排除对角线)
    mask = torch.eye(batch_size, device=z1.device, dtype=torch.bool)
    neg_sim = sim_matrix.masked_fill(mask, -1e9)
    all_sim = torch.logsumexp(neg_sim, dim=1)
    
    return - (pos_sim - all_sim).mean()

def diversity_loss(z):
    """特征熵损失:鼓励特征分布分散,防止过度聚集"""
    # 归一化特征后计算熵
    z_norm = F.normalize(z, dim=1)
    # 计算每个特征维度的熵
    prob = F.softmax(z_norm, dim=1)
    entropy = - (prob * torch.log(prob + 1e-8)).sum(dim=1).mean()
    # 熵越大,特征越分散;乘以小权重平衡损失
    return entropy * 0.05

def total_loss(z1, z2):
    """总损失 = InfoNCE + 多样性约束"""
    contrast_loss = infonce_loss(z1, z2)
    div_loss = diversity_loss(z1) + diversity_loss(z2)
    return contrast_loss + div_loss

# --------------------------
# 5. 带早停的训练流程(150轮+早停)
# --------------------------
def main():
    # 设备设置
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f"使用设备: {device}")
    
    # 1. 加载数据集(全量+大批次)
    dataset = TUDataset(root='./data/MUTAG', name='MUTAG').shuffle()
    train_dataset = dataset[:150]
    # 批次大小设为32(避免图例过多,同时保证负样本数量)
    train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
    
    # 2. 初始化模型
    in_channels = dataset.num_node_features
    hidden_channels = 256
    out_channels = 128
    encoder = EnhancedGCNEncoder(in_channels, hidden_channels, out_channels).to(device)
    
    # 3. 优化器+学习率衰减(更平缓的衰减)
    optimizer = torch.optim.Adam(encoder.parameters(), lr=0.001, weight_decay=1e-5)
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=150, eta_min=5e-5)
    
    # 记录损失
    loss_history = []
    # 早停参数(连续10轮Loss波动<0.02则停止)
    early_stop_patience = 10
    min_loss_fluctuation = 0.02
    
    # 4. 可视化图增强效果
    single_graph = dataset[0].to(device)
    aug1 = graph_augmentation(single_graph)
    aug2 = graph_augmentation(single_graph)
    plot_augmentation_effect(single_graph, aug1, aug2)
    
    # 5. 训练前可视化(清晰图例)
    print("=== 训练前特征可视化(清晰图例版) ===")
    encoder.eval()
    with torch.no_grad():
        for data in train_loader:
            data = data.to(device)
            view1 = graph_augmentation(data)
            view2 = graph_augmentation(data)
            z1 = encoder(view1.x, view1.edge_index, view1.batch)
            z2 = encoder(view2.x, view2.edge_index, view2.batch)
            if z1.shape[0] > 1:
                plot_feature_tsne(z1, z2, "before")
            break  # 只取第一个批次
    
    # 6. 训练(最多150轮+早停)
    print("=== 开始训练(最多150轮,带早停) ===")
    encoder.train()
    max_epochs = 200
    for epoch in range(1, max_epochs+1):
        total_epoch_loss = 0.0
        for data in train_loader:
            data = data.to(device)
            optimizer.zero_grad()
            
            view1 = graph_augmentation(data)
            view2 = graph_augmentation(data)
            
            z1_train = encoder(view1.x, view1.edge_index, view1.batch)
            z2_train = encoder(view2.x, view2.edge_index, view2.batch)
            
            # 计算总损失
            loss = total_loss(z1_train, z2_train)
            loss.backward()
            optimizer.step()
            
            total_epoch_loss += loss.item() * data.num_graphs
        
        # 学习率衰减
        scheduler.step()
        
        # 计算平均损失
        avg_loss = total_epoch_loss / len(train_loader.dataset)
        loss_history.append(avg_loss)
        
        # 早停判断
        if len(loss_history) > early_stop_patience:
            recent_losses = loss_history[-early_stop_patience:]
            loss_fluctuation = max(recent_losses) - min(recent_losses)
            if loss_fluctuation < min_loss_fluctuation:
                print(f"早停触发!Epoch {epoch}, Loss波动 {loss_fluctuation:.4f} < {min_loss_fluctuation}")
                break
        
        # 每15轮打印一次
        if epoch % 15 == 0:
            current_lr = optimizer.param_groups[0]['lr']
            print(f"Epoch: {epoch}, Avg Loss: {avg_loss:.4f}, LR: {current_lr:.6f}")
    
    # 7. 训练后可视化(清晰图例)
    print("=== 训练后特征可视化(清晰图例版) ===")
    encoder.eval()
    with torch.no_grad():
        for data in train_loader:
            data = data.to(device)
            view1 = graph_augmentation(data)
            view2 = graph_augmentation(data)
            z1 = encoder(view1.x, view1.edge_index, view1.batch)
            z2 = encoder(view2.x, view2.edge_index, view2.batch)
            if z1.shape[0] > 1:
                plot_feature_tsne(z1, z2, "after")
            break
    
    # 8. 可视化损失曲线
    plot_loss_curve(loss_history)
    
    # 保存模型
    torch.save(encoder.state_dict(), 'anti_overfit_gcl_encoder_clear_legend.pth')
    print("\n=== 训练完成(清晰图例版) ===")
    print("生成的核心文件:")
    print("- feature_tsne_before_clear_legend.png: 训练前特征(清晰图例)")
    print("- feature_tsne_after_clear_legend.png: 训练后特征(清晰图例)")
    print("- loss_curve_anti_overfit.png: 损失曲线(带早停标记)")
    print("- anti_overfit_gcl_encoder_clear_legend.pth: 最终模型")

if __name__ == "__main__":
    main()

实验结果如下:
在这里插入图片描述
=== 训练前特征可视化===
在这里插入图片描述
=== 开始训练(最多150轮,带早停) ===
Epoch: 15, Avg Loss: 2.9646, LR: 0.000977
Epoch: 30, Avg Loss: 2.9109, LR: 0.000909
Epoch: 45, Avg Loss: 2.7877, LR: 0.000804
Epoch: 60, Avg Loss: 2.8409, LR: 0.000672
Epoch: 75, Avg Loss: 2.7608, LR: 0.000525
Epoch: 90, Avg Loss: 2.6280, LR: 0.000378
Epoch: 105, Avg Loss: 2.6951, LR: 0.000246
Epoch: 120, Avg Loss: 2.7887, LR: 0.000141
Epoch: 135, Avg Loss: 2.6862, LR: 0.000073
Epoch: 150, Avg Loss: 2.7332, LR: 0.000050

=== 训练后特征可视化===
在这里插入图片描述
在这里插入图片描述
未完待整理。。。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐