1. 项目概述:图神经网络与社交关系预测
在当今数据驱动的时代,社交网络分析已成为理解人类行为模式的重要工具。作为一名长期从事机器学习实践的工程师,我发现传统方法在处理社交网络这类图结构数据时往往捉襟见肘。这正是图神经网络(GNN)大显身手的地方——它能够直接处理节点和边的关系数据,完美契合社交网络分析的需求。
本项目将使用PyTorch Geometric(PyG)这一专业的图神经网络框架,构建一个完整的社交关系预测系统。不同于常见的图像或文本处理任务,这里的核心挑战在于如何有效捕捉社交网络中的拓扑结构和节点特征。我们将从零开始,涵盖数据准备、模型构建、训练优化到结果可视化的全流程。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 环境配置与数据准备
2.1 开发环境搭建
工欲善其事,必先利其器。我强烈建议使用conda创建独立的Python环境,避免依赖冲突:
bash复制conda create -n pyg_env python=3.8
conda activate pyg_env
pip install torch torchvision torchaudio
pip install torch-geometric
pip install networkx matplotlib scikit-learn
注意:PyTorch Geometric需要与PyTorch版本严格匹配。如果遇到安装问题,建议参考官方文档选择兼容版本组合。
2.2 构建社交网络图数据
真实的社交网络数据往往包含复杂的属性和关系。为便于理解,我们先构建一个包含100个节点的模拟社交网络:
python复制import torch
from torch_geometric.data import Data
import networkx as nx
import numpy as np
# 生成100个节点的社交网络
num_nodes = 100
num_edges = 300 # 平均每个节点有6条边
# 节点特征:年龄(归一化)、5个兴趣标签(one-hot编码)
age = torch.rand(num_nodes, 1) # 0-1范围
interests = torch.zeros(num_nodes, 5)
interests[torch.arange(num_nodes), torch.randint(0,5,(num_nodes,))] = 1
x = torch.cat([age, interests], dim=1) # 组合成6维特征
# 生成随机边(避免自环)
edge_index = []
for _ in range(num_edges):
while True:
i, j = torch.randint(0, num_nodes, (2,))
if i != j and (i,j) not in edge_index:
edge_index.append([i,j])
break
edge_index = torch.tensor(edge_index).t().contiguous()
# 创建图数据对象
data = Data(x=x, edge_index=edge_index)
print(f"图数据构建完成:{num_nodes}个节点,{edge_index.shape[1]}条边")
这个模拟数据集中,每个节点有6维特征:
- 第1维:归一化的年龄(0-1)
- 后5维:one-hot编码的兴趣标签
3. 图神经网络模型设计
3.1 GCN架构详解
我们采用图卷积网络(GCN)作为基础架构,其核心思想是通过邻域聚合来更新节点表示:
python复制import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import GCNConv
class SocialGNN(nn.Module):
def __init__(self, input_dim, hidden_dim, output_dim):
super(SocialGNN, self).__init__()
self.conv1 = GCNConv(input_dim, hidden_dim)
self.conv2 = GCNConv(hidden_dim, output_dim)
self.dropout = nn.Dropout(0.5)
def forward(self, data):
x, edge_index = data.x, data.edge_index
# 第一层GCN
x = self.conv1(x, edge_index)
x = F.relu(x)
x = self.dropout(x)
# 第二层GCN
x = self.conv2(x, edge_index)
return x
关键设计考虑:
- 使用两层GCN捕捉局部和全局结构信息
- 加入Dropout(0.5)防止过拟合
- ReLU激活引入非线性
3.2 链接预测策略
链接预测需要评估节点对的连接可能性。我们采用以下策略:
python复制def predict_link(embeddings, node_i, node_j):
# 余弦相似度作为连接得分
score = F.cosine_similarity(embeddings[node_i].unsqueeze(0),
embeddings[node_j].unsqueeze(0))
return score.item()
def generate_negative_samples(data,
