使用 PyG 库构建 GCN 网络进行多标签分类任务

本文提供了一个使用 PyG 库构建 GCN 网络实现多标签分类任务的完整代码示例。该示例针对以下数据特点:

  • 图数量: 42 个
  • 节点数量: 37 个
  • 图像尺寸: 40 像素
  • 标签数量: 8 个
  • 边数量: 61 条
  • 节点特征文件: 'C:\Users\jh\Desktop\data\input\images\{i}.png_{j}.png',其中 {i} 代表图的索引,{j} 代表节点的索引。每个节点特征包含图像像素值,对应 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',第一列为源节点,第二列为目标节点,没有表头。

目标: 预测每个节点的 8 维标签向量,并使预测标签向量与真实标签向量一致。

方法: 使用 PyG 库建立 GCN 网络,并利用节点颜色特征进行训练。

代码示例:

import torch
import torch.nn as nn
import torch.nn.functional as F
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 = self.conv2(x, edge_index)
        return x

# 加载节点特征
node_features = []
for i in range(42):  # 图的数量
    for j in range(37):  # 节点数量
        image_path = f'C:\Users\jh\Desktop\data\input\images\{i}.png_{j}.png'
        image = load_image(image_path)  # 加载图像并获取像素值
        node_features.append(image)
node_features = torch.stack(node_features)

# 加载标签
labels = []
for i in range(42):
    for j in range(37):
        label_path = f'C:\Users\jh\Desktop\data\input\labels\{i}_{j}.txt'
        with open(label_path, 'r') as file:
            label = list(map(int, file.read().split()))
        labels.append(label)
labels = torch.tensor(labels, dtype=torch.float)

# 加载边
edges = []
edge_index = []
with open('C:\Users\jh\Desktop\data\input\edges_L.csv', 'r') as file:
    for line in file:
        src, tgt = map(int, line.split(','))
        edges.append((src, tgt))
        edges.append((tgt, src))
        edge_index.append((src, tgt))
        edge_index.append((tgt, src))
edges = torch.tensor(edges, dtype=torch.long).t().contiguous()
edge_index = torch.tensor(edge_index, dtype=torch.long).t().contiguous()

# 分割训练和验证掩码
train_mask = torch.zeros(42 * 37, dtype=torch.bool)
val_mask = torch.zeros(42 * 37, dtype=torch.bool)
train_mask[:30 * 42] = 1
val_mask[30 * 42:] = 1

# 创建 PyG 数据对象
data = Data(x=node_features, edge_index=edge_index, y=labels)

# 创建数据加载器
loader = DataLoader([data], batch_size=1)

# 创建 GCN 模型
model = GCN(num_features=8, num_labels=8)

# 定义损失函数
criterion = nn.BCEWithLogitsLoss()

# 定义优化器
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

# 训练循环
model.train()
for epoch in range(100):
    for batch in loader:
        optimizer.zero_grad()
        out = model(batch.x, batch.edge_index)
        loss = criterion(out[train_mask], batch.y[train_mask])
        loss.backward()
        optimizer.step()

# 验证循环
model.eval()
with torch.no_grad():
    out = model(data.x, data.edge_index)
    pred_labels = torch.sigmoid(out[val_mask]) > 0.5

# 将预测标签转换为 8 维向量
pred_labels = pred_labels.view(42, 37, 8)

# 打印预测标签
for i in range(42):
    for j in range(37):
        print(f'图 {i}, 节点 {j}: {pred_labels[i, j].tolist()}')

注意:

  • load_image 函数需要根据具体的图像加载方法进行实现。
  • 代码中的模型、优化器和损失函数等参数可以根据具体需求进行调整。
  • 训练和验证掩码的设置可以根据实际情况进行调整,本示例将每个图的前 30 个节点用于训练,后 7 个节点用于验证。
  • 输出的预测标签是 8 维的二进制向量,表示每个标签是否被预测为正类。
  • 代码示例中的 num_graphs, num_nodes, num_features, num_labels 等参数需要根据实际数据进行修改。

希望本文提供的代码示例能够帮助您理解如何使用 PyG 库构建 GCN 网络实现多标签分类任务。

使用 PyG 库构建 GCN 网络进行多标签分类任务

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

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