使用PyG库构建GCN网络实现多标签分类任务
使用PyG库构建GCN网络实现多标签分类任务
本代码使用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',标签用空格隔开,例如某个节点的标签向量为: '2 2 1 1 3 1 2 1','5_21.txt'的标签向量为'1 3 4 1 3 1 1 3'
- 真实标签值:只有'0'、'1'、'2'、'3'、'4'五个类别,但每个节点的标签是一个8维的标签向量
- 边关系文件:'C:\Users\jh\Desktop\data\input\edges_L.csv',表格中没有header,第一列为源节点,第二列为目标节点,共有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
from torch_geometric.data import DataLoader
from torch.utils.data.dataset import random_split
import os
import numpy as np
import pandas as pd
from PIL import Image
# 设置超参数
num_graphs = 42
num_nodes = 37
image_size = 40
num_labels = 8
num_edges = 61
num_classes = 5
hidden_channels = 16
lr = 0.01
num_epochs = 100
# 定义GCN模型
class GCN(nn.Module):
def __init__(self, hidden_channels, num_classes):
super(GCN, self).__init__()
self.conv1 = GCNConv(num_labels, hidden_channels)
self.conv2 = GCNConv(hidden_channels, num_classes)
def forward(self, x, edge_index):
x = self.conv1(x, edge_index)
x = F.relu(x)
x = F.dropout(x, training=self.training)
x = self.conv2(x, edge_index)
return x
# 读取节点特征
node_features = []
for i in range(num_graphs):
for j in range(num_nodes):
image_path = f'C:\Users\jh\Desktop\data\input\images\{i}.png_{j}.png'
image = Image.open(image_path)
image = np.array(image)
node_features.append(image)
node_features = torch.tensor(node_features, dtype=torch.float32)
# 读取边的关系
edge_index = []
edges_df = pd.read_csv('C:\Users\jh\Desktop\data\input\edges_L.csv', header=None)
edges = edges_df.values.tolist()
for edge in edges:
edge_index.append(edge)
edge_index.append([edge[1], edge[0]]) # 无向边,需要添加反向边
edge_index = torch.tensor(edge_index, dtype=torch.long).t().contiguous()
# 读取标签向量
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 = f.read().split()
label = list(map(int, label))
labels.append(label)
labels = torch.tensor(labels, dtype=torch.float32)
# 划分训练集和验证集
train_mask = torch.zeros(num_graphs * num_nodes, dtype=torch.uint8)
train_mask[:30*num_nodes] = 1 # 前30个节点加入训练掩码
val_mask = torch.zeros(num_graphs * num_nodes, dtype=torch.uint8)
val_mask[30*num_nodes:] = 1 # 后7个节点加入验证掩码
# 构建图数据
data = Data(x=node_features, edge_index=edge_index, y=labels, train_mask=train_mask, val_mask=val_mask)
# 划分训练集和验证集
dataset = [data]
train_dataset, val_dataset = random_split(dataset, [30*num_nodes*num_graphs, 7*num_nodes*num_graphs])
# 创建数据加载器
train_loader = DataLoader(train_dataset, batch_size=num_nodes*num_graphs, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=num_nodes*num_graphs, shuffle=False)
# 初始化模型
model = GCN(hidden_channels, num_classes)
# 定义损失函数和优化器
criterion = nn.BCEWithLogitsLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
# 训练模型
for epoch in range(num_epochs):
model.train()
for data in train_loader:
out = model(data.x, data.edge_index)
loss = criterion(out[data.train_mask], data.y[data.train_mask])
optimizer.zero_grad()
loss.backward()
optimizer.step()
model.eval()
for data in val_loader:
with torch.no_grad():
out = model(data.x, data.edge_index)
pred = out[data.val_mask]
pred[pred >= 0.5] = 1
pred[pred < 0.5] = 0
accuracy = (pred == data.y[data.val_mask]).sum().item() / pred.size(0)
print(f'Epoch: {epoch+1}, Loss: {loss.item()}, Validation Accuracy: {accuracy}')
# 预测节点特征向量和标签向量
model.eval()
for data in val_loader:
with torch.no_grad():
out = model(data.x, data.edge_index)
pred = out[data.val_mask]
pred[pred >= 0.5] = 1
pred[pred < 0.5] = 0
print('Predicted Labels:')
print(pred)
print('True Labels:')
print(data.y[data.val_mask])
请确保将节点特征、边的关系和标签向量的文件路径替换成您实际的文件路径。此代码将输出每个节点的预测标签向量,并与真实标签向量进行比较。
原文地址: https://www.cveoy.top/t/topic/pl1I 著作权归作者所有。请勿转载和采集!