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

于是这里简单整理一下对比学习的初步实践,后续有时间再整理对比学习的前沿研究。
什么是对比学习?
对比学习(Contrastive Learning)是一类自监督 / 弱监督表征学习方法,核心是通过区分正负样本对,让模型在特征空间中拉近相似样本、推远不相似样本,从而学到数据的内在结构与鲁棒表示,大幅减少对人工标注的依赖。它在计算机视觉、自然语言处理等领域应用广泛,是解决数据稀缺问题的关键技术之一。
- 核心逻辑:模型不直接学习标签映射,而是通过成对样本的相似性判断来优化特征表示 —— 对同一实例的不同 “视角”(正样本对)拉近距离,对不同实例(负样本对)推远距离。
- 学习目标:在嵌入空间中,最大化正样本对的相似度、最小化负样本对的相似度,使特征具备更强的区分力与泛化能力。

Question 1:在对比学习中,什么才算是正样本或者负样本?不一定是做分类任务才用到对比学习吧?
A1:正负样本的定义和任务类型没有必然绑定。
- 对比学习不只是用于分类任务,它的核心是通过样本间的相似性关系学习通用表征,适配检索、匹配、聚类等多种任务;
- 正负样本的划分,本质是基于 “语义或实例层面的相似性”,而非标签 —— 不同任务场景下,判断相似性的标准完全不同。
Question 2:对比学习最早源自于哪篇论文?研究是如何发展的?
A2:核心对比学习框架的奠基性论文如下:
- Sumit Chopra、Raia Hadsell 和 Yann LeCun 发表于 CVPR 2005 的 《Dimensionality Reduction by Learning an Invariant Mapping》,首次系统性提出对比学习的核心框架与对比损失函数(Contrastive Loss),通过学习数据的不变映射,拉近相似样本、拉远不相似样本,成为后续对比学习的核心理论基础。该论文将对比学习用于降维任务,为自监督表示学习提供了重要范式。(论文解读)
- 《Contrastive Predictive Coding》(Aaron van den Oord 等,2018):提出对比预测编码(CPC),通过预测未来样本的上下文表示,实现序列数据的无监督表示学习,为现代对比学习在语音、文本等序列数据领域的应用奠定了基础。(论文解读链接)
- 《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
=== 训练后特征可视化===


未完待整理。。。
更多推荐
所有评论(0)