PyTorch如何将可变长度图数据加载到DataLoader供GCN训练
报错原因
你遇到的维度报错核心是用法不符合PyTorch Geometric(PyG)的图数据规范:
- 图分类任务的每个样本是独立图,节点数、边数都可能不同,不能把多个图的边、节点特征直接拼接成单个tensor构造全局
Data对象——你两个样本分别有5条边、3条边,硬拼到一个tensor里自然会出现维度不匹配 - 原代码里
y=y属于未定义变量的笔误 - 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
相关产品推荐
相关产品推荐

