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 个节点颜色特征加入验证掩码。
模型实现
使用 PyG 库建立 GCN 网络实现多标签分类任务。
import os
import torch
import torch.nn.functional as F
from torch_geometric.data import InMemoryDataset, DataLoader
from torch_geometric.nn import GCNConv
from torch_geometric.utils import from_networkx
import pandas as pd
import numpy as np
from PIL import Image
class GraphDataset(InMemoryDataset):
def __init__(self, root, transform=None, pre_transform=None):
super(GraphDataset, self).__init__(root, transform, pre_transform)
self.data, self.slices = torch.load(self.processed_paths[0])
@property
def raw_file_names(self):
return ['edges_L.csv']
@property
def processed_file_names(self):
return ['data.pt']
def download(self):
# Download the dataset from the given url
pass
def process(self):
# Read edge data from csv file
edge_data = pd.read_csv(self.raw_paths[0], header=None)
edge_index = torch.tensor(edge_data.values, dtype=torch.long).t().contiguous()
# Read node features and labels
num_nodes = 37
num_labels = 8
x = torch.zeros(num_nodes, 40) # Placeholder for node features
y = torch.zeros(num_nodes, num_labels) # Placeholder for node labels
for i in range(num_nodes):
# Read node features
image_path = f'C:\Users\jh\Desktop\data\input\images\{i}.png_{j}.png'
image = Image.open(image_path).convert('RGB')
image = transform(image) # Apply any necessary transformations
x[i] = image.view(-1, 40)
# Read node labels
label_path = f'C:\Users\jh\Desktop\data\input\labels{i}{j}.txt'
with open(label_path, 'r') as file:
labels = [int(label) for label in file.readline().split()]
y[i] = torch.tensor(labels)
data = from_networkx(edge_index)
data.x = x
data.y = y
# Split the data into training and validation sets
train_mask = torch.zeros(num_nodes, dtype=torch.bool)
train_mask[:30] = 1 # First 30 nodes for training
val_mask = torch.zeros(num_nodes, dtype=torch.bool)
val_mask[30:] = 1 # Last 7 nodes for validation
data.train_mask = train_mask
data.val_mask = val_mask
# Save the preprocessed data
torch.save(self.collate([data]), self.processed_paths[0])
class GCN(torch.nn.Module):
def __init__(self, in_channels, out_channels):
super(GCN, self).__init__()
self.conv1 = GCNConv(in_channels, 16)
self.conv2 = GCNConv(16, out_channels)
def forward(self, x, edge_index):
x = self.conv1(x, edge_index)
x = F.relu(x)
x = self.conv2(x, edge_index)
return x
# Initialize the dataset
dataset = GraphDataset(root='data/')
# Create the dataloader
loader = DataLoader(dataset, batch_size=1, shuffle=False)
# Initialize the model
model = GCN(in_channels=40, out_channels=8)
# Set device (CPU or GPU)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# Move model to the device
model = model.to(device)
# Set optimizer and loss function
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
criterion = torch.nn.CrossEntropyLoss()
# Training loop
model.train()
for data in loader:
data = data.to(device)
optimizer.zero_grad()
out = model(data.x, data.edge_index)
loss = criterion(out[data.train_mask], data.y[data.train_mask].max(dim=1)[1])
loss.backward()
optimizer.step()
# Evaluation loop
model.eval()
for data in loader:
data = data.to(device)
out = model(data.x, data.edge_index)
pred = out.argmax(dim=1)
print(pred)
代码解释
- GraphDataset 类
- 读取边的信息并构建邻接矩阵(edge_index)。
- 读取每个节点的图像特征和标签。
- 将数据分割成训练集和验证集。
- 将处理后的数据保存为 data.pt 文件。
- GCN 类
- 定义 GCN 网络结构,包括两层 GCNConv 层。
- 使用 ReLU 激活函数。
- 训练和评估
- 初始化模型、优化器和损失函数。
- 使用 DataLoader 迭代训练数据,进行模型训练。
- 在训练完成后,使用 DataLoader 迭代验证数据,进行模型评估。
注意
- 代码中的 transform 函数需要根据你的实际需求进行定义和实现,用于对图像数据进行预处理。
- 代码中使用 argmax 函数获取预测标签,假设真实标签为 one-hot 编码,并使用 CrossEntropyLoss 作为损失函数。
- 训练过程需要根据实际情况调整超参数,例如学习率、batch size 等。
完整代码
# ...
请注意,此处的代码仅为一个示例,需要根据实际情况进行相应的修改和调整。
原文地址: https://www.cveoy.top/t/topic/pl2e 著作权归作者所有。请勿转载和采集!