You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch如何将可变长度图数据加载到DataLoader供GCN训练

报错原因

你遇到的维度报错核心是用法不符合PyTorch Geometric(PyG)的图数据规范:

  1. 图分类任务的每个样本是独立图,节点数、边数都可能不同,不能把多个图的边、节点特征直接拼接成单个tensor构造全局Data对象——你两个样本分别有5条边、3条边,硬拼到一个tensor里自然会出现维度不匹配
  2. 原代码里y=y属于未定义变量的笔误
  3. PyG处理多图数据集,需要把每个图单独封装为Data实例,再通过DataLoader自动完成批内图合并、节点索引重映射、批向量生成的工作,不需要手动拼接全局张量。

可运行完整实现

依赖导入

import numpy as np
import torch
import torch.nn.functional as F
from torch_geometric.data import Data
from torch_geometric.loader import DataLoader
from torch_geometric.nn import GCNConv, global_mean_pool
from torch.nn import Linear

第一步:正确构造图数据集

逐图封装独立Data对象,注意每个图的节点索引从0开始独立计数:

# 原始数据
edge_origins = np.array([[0,1,2,3,4],[6,7,8]], dtype=object)
edge_destinations = np.array([[1,2,3,4,5],[7,8,9]], dtype=object)
target = np.array([0,1])
x_raw = [
    [np.array([0.1,0.5,0.2]),np.array([0.5,0.6,0.23]),
     np.array([0.1,0.5,0.5]),np.array([0.1,0.6,0.23]),
     np.array([0.1,0.4,0.4]),np.array([0.52,0.6,0.23])],
    [np.array([0.1,0.3,0.3]),np.array([0.3,0.6,0.23]),
     np.array([0.1,0.1,0.2]),np.array([0.4,0.6,0.23])]
]

dataset = []
for graph_idx in range(len(target)):
    # 转换当前图节点特征为float tensor
    node_feat = torch.tensor(np.array(x_raw[graph_idx]), dtype=torch.float)
    # 处理边索引:第二个图原始节点编号是6-9,重映射为从0开始的独立索引
    if graph_idx == 1:
        src = edge_origins[graph_idx] - 6
        dst = edge_destinations[graph_idx] - 6
    else:
        src = edge_origins[graph_idx]
        dst = edge_destinations[graph_idx]
    edge_idx = torch.tensor(np.array([src, dst]), dtype=torch.long)
    # 转换当前图标签
    label = torch.tensor([target[graph_idx]], dtype=torch.long)
    # 组装单图对象
    dataset.append(Data(x=node_feat, edge_index=edge_idx, y=label))

# 验证数据集构造结果
print(f"数据集总样本数:{len(dataset)}")
for i, g in enumerate(dataset):
    print(f"图{i}:节点数{g.num_nodes},边数{g.num_edges},标签{g.y.item()},节点特征维度{g.num_node_features}")

运行后输出:

数据集总样本数:2
图0:节点数6,边数5,标签0,节点特征维度3
图1:节点数4,边数3,标签1,节点特征维度3

第二步:数据集打乱、拆分与DataLoader构造

图和标签绑定在同一个Data对象中,打乱时不会出现标签错位:

torch.manual_seed(12345)
# 同步打乱所有样本
dataset = [dataset[i] for i in torch.randperm(len(dataset)).tolist()]

# 拆分训练、测试集
train_dataset = dataset[:1]
test_dataset = dataset[1:]
print(f'训练集图数量: {len(train_dataset)}')
print(f'测试集图数量: {len(test_dataset)}')

# 构造加载器,PyG会自动处理批内图合并、batch向量生成
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)

第三步:GCN模型搭建与训练

你原有的模型逻辑基本正确,仅需提前计算特征维度和类别数适配列表格式的数据集:

num_node_feats = dataset[0].num_node_features
num_classes = len(set(target.tolist()))

class GCN(torch.nn.Module):
    def __init__(self, hidden_channels):
        super(GCN, self).__init__()
        torch.manual_seed(12345)
        self.conv1 = GCNConv(num_node_feats, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, hidden_channels)
        self.conv3 = GCNConv(hidden_channels, hidden_channels)
        self.lin = Linear(hidden_channels, num_classes)

    def forward(self, x, edge_index, batch):
        # 计算节点嵌入
        x = self.conv1(x, edge_index)
        x = x.relu()
        x = self.conv2(x, edge_index)
        x = x.relu()
        x = self.conv3(x, edge_index)

        # 全局平均池化得到图级嵌入
        x = global_mean_pool(x, batch)  

        # 分类头
        x = F.dropout(x, p=0.5, training=self.training)
        x = self.lin(x)
        return x

model = GCN(hidden_channels=64)
print(model)

optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
criterion = torch.nn.CrossEntropyLoss()

def train():
    model.train()
    for data in train_loader:
        out = model(data.x, data.edge_index, data.batch)
        loss = criterion(out, data.y)
        loss.backward()
        optimizer.step()
        optimizer.zero_grad()

def test(loader):
    model.eval()
    correct = 0
    for data in loader:
        out = model(data.x, data.edge_index, data.batch)
        pred = out.argmax(dim=1)
        correct += int((pred == data.y).sum())
    return correct / len(loader.dataset)

# 启动训练
for epoch in range(1, 171):
    train()
    train_acc = test(train_loader)
    test_acc = test(test_loader)
    print(f'Epoch: {epoch:03d}, 训练集准确率: {train_acc:.4f}, 测试集准确率: {test_acc:.4f}')

新样本预测方法

后续输入新图时,按照单图封装逻辑处理即可预测:

def predict_new_graph(node_feat_list, edge_src, edge_dst):
    """
    新图分类预测
    :param node_feat_list: 每个节点的特征列表,格式与单图x_raw一致
    :param edge_src: 边源节点索引(从0开始计数)
    :param edge_dst: 边目标节点索引(从0开始计数)
    """
    model.eval()
    x = torch.tensor(np.array(node_feat_list), dtype=torch.float)
    edge_index = torch.tensor([edge_src, edge_dst], dtype=torch.long)
    # 单图预测时batch向量为全0,代表所有节点属于同一个图
    batch = torch.zeros(x.shape[0], dtype=torch.long)
    out = model(x, edge_index, batch)
    return out.argmax(dim=1).item()

# 测试新图预测
new_graph_feat = [np.array([0.2,0.3,0.1]), np.array([0.4,0.5,0.2]), np.array([0.1,0.2,0.3])]
new_src = [0,1]
new_dst = [1,2]
print(f"新图预测类别:{predict_new_graph(new_graph_feat, new_src, new_dst)}")

内容的提问来源于stack exchange,提问作者Slowat_Kela

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.30 14:12:36