使用 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
多标签图神经网络 (GCN) 预测节点标签 - PYG 实现

原文地址: https://www.cveoy.top/t/topic/pl2o 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录