PyTorch Geometric (PYG) GCN网络实现多标签分类任务 - 完整代码示例
使用PyTorch Geometric (PYG) 库构建GCN网络的完整代码示例
本代码示例展示了如何使用PyTorch Geometric (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',每个文件代表一个节点的图像,其像素值作为节点特征。 - 标签向量文件:
'C:\Users\jh\Desktop\data\input\labels\{i}_{j}.txt',每个文件包含一个节点的8维标签向量,标签值用空格隔开。例如,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',第一列为源节点,第二列为目标节点,共有 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
# 设置设备
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 定义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 = self.conv2(x, edge_index)
return x
# 加载节点特征
def load_node_features(image_size, num_nodes):
features = []
for i in range(num_nodes):
img_path = f"C:/Users/jh/Desktop/data/input/images/{i}.png"
img = Image.open(img_path).resize((image_size, image_size))
img_tensor = transforms.ToTensor()(img)
features.append(img_tensor)
return torch.stack(features)
# 加载标签向量
def load_labels(num_nodes):
labels = []
for i in range(num_nodes):
label_path = f"C:/Users/jh/Desktop/data/input/labels/{i}.txt"
with open(label_path, 'r') as file:
label_vector = [int(label) for label in file.read().split()]
labels.append(label_vector)
return torch.tensor(labels)
# 加载边的关系
def load_edges(num_edges):
edges = []
with open("C:/Users/jh/Desktop/data/input/edges_L.csv", 'r') as file:
for line in file.readlines():
source, target = line.strip().split(',')
edges.append((int(source), int(target)))
return torch.tensor(edges).t().contiguous()
# 构建图数据
def build_graph_data(num_graphs, num_nodes, image_size, num_labels, num_edges):
node_features = load_node_features(image_size, num_nodes)
labels = load_labels(num_nodes)
edges = load_edges(num_edges)
train_mask = torch.zeros(num_nodes, dtype=torch.bool)
train_mask[:30] = 1 # 前30个节点作为训练节点
val_mask = torch.zeros(num_nodes, dtype=torch.bool)
val_mask[30:] = 1 # 后7个节点作为验证节点
data = Data(x=node_features, edge_index=edges, y=labels)
data.train_mask = train_mask
data.val_mask = val_mask
return data
# 训练模型
def train_model(model, data, epochs, lr):
model = model.to(device)
data = data.to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
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_label = (pred[data.val_mask] > 0.5).float()
accuracy = (pred_label == data.y[data.val_mask]).float().mean()
print(f'Epoch: {epoch+1}/{epochs}, Loss: {loss.item()}, Accuracy: {accuracy.item()}')
# 主函数
if __name__ == '__main__':
num_graphs = 42
num_nodes = 37
image_size = 40
num_labels = 8
num_edges = 61
epochs = 100
lr = 0.01
data = build_graph_data(num_graphs, num_nodes, image_size, num_labels, num_edges)
model = GCN(num_features=3, num_labels=num_labels)
train_model(model, data, epochs, lr)
注意:
- 上述代码仅提供了GCN网络的基本结构和训练流程,其中的图像处理部分(如图像加载、像素值转换等)需要您根据实际情况进行适当修改。
- 您可能还需要安装必要的依赖库,如PyTorch Geometric、PIL等。
代码解析:
- 定义GCN模型:
GCN类定义了 GCN 网络,包含两个 GCNConv 层,使用 ReLU 激活函数。 - 加载节点特征:
load_node_features函数加载图像文件并将其转换为张量,每个张量代表一个节点的特征。 - 加载标签向量:
load_labels函数加载标签文件,将每个文件的内容转换为标签向量。 - 加载边关系:
load_edges函数加载边关系文件,并将其转换为张量。 - 构建图数据:
build_graph_data函数将节点特征、标签向量和边关系组合成一个Data对象,并设置训练掩码和验证掩码。 - 训练模型:
train_model函数使用Adam优化器训练 GCN 模型,并计算训练损失和验证精度。
进一步优化:
- 可以根据实际需要调整模型结构、训练参数等。
- 可以使用其他损失函数,例如多标签交叉熵损失。
- 可以使用不同的优化器,例如 SGD、RMSprop 等。
- 可以添加正则化、dropout 等技术来防止过拟合。
原文地址: https://www.cveoy.top/t/topic/pl2w 著作权归作者所有。请勿转载和采集!