使用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'csv文件中,表格中没有header,第一列为源节点,第二列为目标节点,共有61条无向边。
要求输出每个节点的预测特征向量,并根据这些预测特征得到预测标签向量,使预测标签向量与真实标签向量一致。每个节点的预测标签都是一个8维向量,而不是输出概率向量。将每个图的前30个节点颜色特征加入训练掩码,后7个节点颜色特征加入验证掩码。
代码实现
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch_geometric.data import Data, DataLoader
from torch_geometric.nn import GCNConv
# 定义GCN模型
class GCN(nn.Module):
def __init__(self, num_features, num_labels):
super(GCN, self).__init__()
self.conv1 = GCNConv(num_features, 16)
self.conv2 = GCNConv(16, num_labels)
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
# 加载节点特征和标签
def load_data():
x = []
y = []
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 = load_image(image_path)
x.append(image)
# 加载标签
label_path = f'C:\Users\jh\Desktop\data\input\labels{i}{j}.txt'
label = load_label(label_path)
y.append(label)
x = torch.tensor(x, dtype=torch.float)
y = torch.tensor(y, dtype=torch.float)
return x, y
# 加载边关系
def load_edges():
edges = []
with open('C:\Users\jh\Desktop\data\input\edges_L.csv', 'r') as file:
for line in file:
source, target = line.strip().split(',')
edges.append((int(source), int(target)))
edges.append((int(target), int(source)))
edge_index = torch.tensor(edges, dtype=torch.long).t().contiguous()
return edge_index
# 加载图数据
def load_graph_data():
x, y = load_data()
edge_index = load_edges()
data = Data(x=x, edge_index=edge_index, y=y)
return data
# 加载训练和验证数据
def load_train_val_data(data):
train_mask = torch.zeros(num_graphs * num_nodes, dtype=torch.bool)
val_mask = torch.zeros(num_graphs * num_nodes, dtype=torch.bool)
train_mask[:30*num_nodes] = 1
val_mask[30*num_nodes:] = 1
data.train_mask = train_mask
data.val_mask = val_mask
return data
# 训练模型
def train(model, data, epochs):
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.to(device)
data = data.to(device)
optimizer = optim.Adam(model.parameters(), lr=0.01)
for epoch in range(epochs):
model.train()
optimizer.zero_grad()
out = model(data.x, data.edge_index)
loss = F.binary_cross_entropy_with_logits(out[data.train_mask], data.y[data.train_mask])
loss.backward()
optimizer.step()
model.eval()
with torch.no_grad():
pred = model(data.x, data.edge_index)
pred = torch.sigmoid(pred)
pred = torch.round(pred)
train_acc = (pred[data.train_mask] == data.y[data.train_mask]).sum().item() / data.train_mask.sum().item()
val_acc = (pred[data.val_mask] == data.y[data.val_mask]).sum().item() / data.val_mask.sum().item()
print(f'Epoch: {epoch+1}, Train Acc: {train_acc:.4f}, Val Acc: {val_acc:.4f}')
# 加载图片特征
def load_image(image_path):
# 加载图片像素值,将其转换为特征向量
# 这里需要根据具体的图片处理方式进行实现
return image_feature
# 加载标签
def load_label(label_path):
with open(label_path, 'r') as file:
label = [int(x) for x in file.readline().strip().split(' ')]
return label
# 设置参数
num_graphs = 42
num_nodes = 37
image_size = 40
num_labels = 8
num_edges = 61
# 加载图数据
data = load_graph_data()
# 加载训练和验证数据
data = load_train_val_data(data)
# 创建GCN模型
model = GCN(image_size, num_labels)
# 训练模型
train(model, data, epochs=100)
代码解释
- 数据加载:
load_data()函数加载节点特征和标签,load_edges()函数加载边关系,load_graph_data()函数整合数据为Data对象。 - 训练和验证数据:
load_train_val_data()函数将数据分为训练集和验证集,并设置相应的掩码。 - 模型定义:
GCN类定义了GCN模型架构,包含两层GCN卷积层。 - 训练:
train()函数进行模型训练,使用Adam优化器,计算二元交叉熵损失,并输出训练集和验证集的准确率。 - 图片特征和标签加载:
load_image()函数加载图片特征,load_label()函数加载标签,需要根据实际情况进行实现。
注意: 该代码示例仅供参考,需要根据实际任务进行调整。例如,模型架构、优化器和训练参数等都需要根据具体情况进行调整。
下一步: 可以根据实际任务需求,对代码进行进一步完善和优化,例如:
- 添加模型评估指标,例如精确率、召回率等
- 使用不同的数据增强方法,提高模型性能
- 使用更复杂的模型架构,提升模型能力
- 优化模型训练参数,例如学习率、批次大小等
原文地址: https://www.cveoy.top/t/topic/pl2s 著作权归作者所有。请勿转载和采集!