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

如何重塑数据以在PyTorch中训练GCN?解决索引匹配错误

解决GCN训练时的IndexError问题

错误原因

你使用普通的torch.utils.data.DataLoader处理图数据时,会直接把每个样本的edge_index(形状[2,3])按batch维度拼接成[2, 3*32],但模型期望的edge_index要对应批量后的节点特征x(形状[32,9])。而原样本的edge_index使用的是全局节点ID,没有针对batch做ID重映射,导致索引时形状不匹配,触发IndexError。另外,普通DataLoader无法正确处理图数据的批量拼接逻辑,必须使用PyTorch Geometric的专用工具类。

解决方案

核心调整点

  1. 改用PyTorch Geometric的Data类封装单样本图数据,统一管理节点特征、标签、边索引和边权重
  2. 使用PyG的DataLoader替代普通DataLoader,它会自动完成多图拼接时的节点ID重映射,避免ID冲突
  3. 适配新的数据格式调整训练逻辑

修改后的完整代码

数据集定义代码

import torch
from torch.utils.data import Dataset
from torch_geometric.data import Data

class S_Dataset(Dataset):
    def __init__(self, df, transform=None):
        self.df = df
        self.transform = transform

    def __len__(self):
        return len(self.df)

    def __getitem__(self, idx):
        row = self.df.iloc[idx]

        # 单个节点特征,形状[9]
        x = torch.tensor([
            row.date.to_pydatetime().timestamp(), 
            row.s1, row.s2, row.s3, row.s4, 
            row.temp, row.rh, row.Location, row.Node
        ], dtype=torch.float)
        # 单个节点标签,形状[1]
        y = torch.tensor([row.Location], dtype=torch.long)
       
        # 3条边的权重,形状[3]
        edge_attr = torch.tensor([
            row.neighbor1_distance, 
            row.neighbor2_distance, 
            row.neighbor3_distance
        ], dtype=torch.float)

        # 边索引,转置后形状[2,3]
        edge_index = torch.tensor([
            [row.Location, row.neighbor1_name],
            [row.Location, row.neighbor2_name],
            [row.Location, row.neighbor3_name]
        ], dtype=torch.long).t()

        # 用PyG的Data对象封装所有数据
        data = Data(x=x.unsqueeze(0), y=y, edge_index=edge_index, edge_attr=edge_attr)

        if self.transform:
            data = self.transform(data)

        return data

Process_Data = S_Dataset(df)

数据集拆分与加载代码

from torch_geometric.loader import DataLoader

train_size = int(len(Process_Data) * 0.8)
test_size = len(Process_Data) - train_size
train_dataset, test_dataset = torch.utils.data.random_split(Process_Data, [train_size, test_size])

# 使用PyG专用的DataLoader处理图数据批量
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=32, shuffle=True)

GCN模型定义代码

import torch
import torch.nn as nn
import torch.optim as optim
from torch_geometric.nn import GCNConv

class Net(nn.Module):
    def __init__(self, num_classes):
        super(Net, self).__init__()
        self.conv1 = GCNConv(9, 128)
        self.conv2 = GCNConv(128, 64)
        self.fc1 = nn.Linear(64, 32)
        self.fc2 = nn.Linear(32, num_classes)

    def forward(self, x, edge_index, edge_attr):
        x = self.conv1(x, edge_index, edge_attr)
        x = torch.relu(x)
        x = self.conv2(x, edge_index, edge_attr)
        x = torch.relu(x)
        x = self.fc1(x)
        x = torch.relu(x)
        x = self.fc2(x)
        return x

# 初始化模型时传入类别数量
model = Net(num_classes=len(location_to_id))

模型训练代码

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

for epoch in range(100):
    total_loss = 0
    model.train()
    for batch in train_loader:
        optimizer.zero_grad()
        # 从PyG的Batch对象中提取字段
        y_pred = model(batch.x, batch.edge_index, batch.edge_attr)
        # 将标签压缩为一维,匹配损失函数要求的输入形状
        loss = criterion(y_pred, batch.y.squeeze())
        loss.backward()
        optimizer.step()
        total_loss += loss.item()
    print(f'Epoch: {epoch} Loss: {total_loss / len(train_loader):.4f}')

关键说明

  • PyG的DataLoader会自动将多个Data对象拼接为Batch对象:
    • batch.x:所有节点特征拼接成[总节点数, 特征维度]
    • batch.edge_index:所有边索引拼接,同时自动给每个子图的节点ID加上偏移量,避免跨图ID冲突
    • batch.edge_attr:所有边权重拼接成[总边数]
    • batch.y:所有标签拼接成[总节点数]
  • 模型forward方法无需额外调整张量形状,x已经是符合要求的二维张量
  • 损失计算时需将batch.y压缩为一维,与预测结果y_pred的形状匹配

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 22:45:59