使用PyTorch Geometric (PYG) 库构建GCN网络的完整代码示例

本代码示例展示了如何使用PyTorch Geometric (PYG) 库构建GCN网络,并用于处理多标签分类任务。

数据设置:

  • 已知:num_graphs = 42num_nodes = 37image_size = 40num_labels = 8num_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等。

代码解析:

  1. 定义GCN模型: GCN 类定义了 GCN 网络,包含两个 GCNConv 层,使用 ReLU 激活函数。
  2. 加载节点特征: load_node_features 函数加载图像文件并将其转换为张量,每个张量代表一个节点的特征。
  3. 加载标签向量: load_labels 函数加载标签文件,将每个文件的内容转换为标签向量。
  4. 加载边关系: load_edges 函数加载边关系文件,并将其转换为张量。
  5. 构建图数据: build_graph_data 函数将节点特征、标签向量和边关系组合成一个 Data 对象,并设置训练掩码和验证掩码。
  6. 训练模型: train_model 函数使用 Adam 优化器训练 GCN 模型,并计算训练损失和验证精度。

进一步优化:

  • 可以根据实际需要调整模型结构、训练参数等。
  • 可以使用其他损失函数,例如多标签交叉熵损失。
  • 可以使用不同的优化器,例如 SGD、RMSprop 等。
  • 可以添加正则化、dropout 等技术来防止过拟合。

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

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