多标签图神经网络 (GCN) 预测节点标签 - PYG 实现
使用 PYG 库建立 GCN 网络实现多标签分类任务
已知条件:
- num_graphs = 42
- num_nodes = 37
- image_size = 40
- num_labels = 8
- num_edges = 61
- 节点特征文件为 'C:\Users\jh\Desktop\data\input\images{i}.png_{j}.png',每个图片代表一个节点的像素值,共8个维度
- 节点标签文件为 'C:\Users\jh\Desktop\data\input\labels{i}{j}.txt',每个文件包含一个节点的 8 维标签向量,用空格隔开,例如 '2 2 1 1 3 1 2 1'。
- 真实标签值只有 0、1、2、3、4 五个类别
- 边关系储存在 'C:\Users\jh\Desktop\data\input\edges_L.csv' 文件中,第一列为源节点,第二列为目标节点,共有 61 条无向边。
任务目标:
- 预测每个节点的 8 维标签向量,使其与真实标签向量一致
- 将每个图的前 30 个节点颜色特征加入训练掩码,后 7 个节点颜色特征加入验证掩码
代码实现:
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.data import Data
from torch_geometric.nn import GCNConv
# 定义 GCN 模型
class GCN(nn.Module):
def __init__(self, input_dim, hidden_dim, output_dim):
super(GCN, self).__init__()
self.conv1 = GCNConv(input_dim, hidden_dim)
self.conv2 = GCNConv(hidden_dim, output_dim)
def forward(self, data):
x, edge_index = data.x, data.edge_index
x = self.conv1(x, edge_index)
x = F.relu(x)
x = self.conv2(x, edge_index)
return x
# 加载节点特征
node_features = []
for i in range(num_graphs):
for j in range(num_nodes):
img_path = f'C:\Users\jh\Desktop\data\input\images\{i}.png_{j}.png'
img = load_image(img_path)
node_features.append(img)
node_features = torch.tensor(node_features, dtype=torch.float)
# 加载标签
labels = []
for i in range(num_graphs):
for j in range(num_nodes):
label_path = f'C:\Users\jh\Desktop\data\input\labels\{i}{j}.txt'
with open(label_path, 'r') as f:
label = [int(x) for x in f.read().split()]
labels.append(label)
labels = torch.tensor(labels, dtype=torch.float)
# 加载边关系
edges = []
with open('C:\Users\jh\Desktop\data\input\edges_L.csv', 'r') as f:
for line in f:
src, tgt = [int(x) for x in line.strip().split(',')]
edges.append((src, tgt))
edges = torch.tensor(edges, dtype=torch.long).t().contiguous()
# 定义训练和验证掩码
train_mask = torch.zeros(num_graphs * num_nodes, dtype=torch.bool)
train_mask[:num_graphs * 30] = 1
val_mask = torch.zeros(num_graphs * num_nodes, dtype=torch.bool)
val_mask[num_graphs * 30:] = 1
# 创建图数据
data = Data(x=node_features, edge_index=edges)
data.y = labels
# 分割训练和验证数据
data.train_mask = train_mask
data.val_mask = val_mask
# 创建 GCN 模型
model = GCN(input_dim=8, hidden_dim=16, output_dim=8)
# 定义损失函数
criterion = nn.BCEWithLogitsLoss()
# 定义优化器
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
# 训练循环
def train():
model.train()
optimizer.zero_grad()
out = model(data)
loss = criterion(out[data.train_mask], data.y[data.train_mask])
loss.backward()
optimizer.step()
# 验证循环
def test():
model.eval()
out = model(data)
pred = (out > 0.5).float()
correct = pred[data.val_mask] == data.y[data.val_mask]
accuracy = correct.sum().item() / data.val_mask.sum().item()
return accuracy
# 训练模型
for epoch in range(100):
train()
accuracy = test()
print(f'Epoch: {epoch+1}, Accuracy: {accuracy}')
注意:
- 代码中假设您已经实现了
load_image函数来加载图像文件,需要根据您的实际情况进行调整。 - 此代码示例仅供参考,您可能需要根据您的具体需求进行修改和调整。
- 在实际应用中,您可能需要调整超参数、损失函数、优化器等,以获得最佳的模型性能。
- 本示例主要展示了如何使用 PYG 库建立 GCN 网络,并通过训练和验证循环来评估模型的性能。
- 您可以根据自己的实际情况,修改代码中的数据路径、模型参数、训练和验证策略等。
其他资源:
- PYG 官方文档:https://pytorch-geometric.readthedocs.io/en/latest/
- GCN 论文:https://arxiv.org/abs/1609.02907
- 多标签分类:https://en.wikipedia.org/wiki/Multi-label_classification
原文地址: https://www.cveoy.top/t/topic/pl2o 著作权归作者所有。请勿转载和采集!