PyG自定义数据集实现GCN时维度不匹配问题求助
解决PyTorch Geometric中edge_label维度不匹配问题
你的问题核心是自定义数据集的edge_label多了一个冗余维度([100836,1] vs 原始的[100836]),导致损失计算时维度不匹配,直接用下面几种方法就能解决:
方法一:加载数据后直接压缩维度
这是最直接的解决方案,在拿到数据集对象后,对edge_label执行维度压缩操作:
- 使用
squeeze()方法(指定维度避免误删其他维度):# 加载自定义数据集后执行 data.edge_label = data.edge_label.squeeze(dim=1) - 或者用
flatten()方法直接展平为一维:data.edge_label = data.edge_label.flatten()
执行后edge_label的维度就会变成[100836],和原始数据集结构一致。
方法二:在CSV加载阶段处理
如果你是用PyG的load_csv_edge_dataset或相关工具加载数据,可以在读取时就处理标签维度:
from torch_geometric.loader import load_csv_edge_dataset # 加载数据时直接对edge_label做处理 dataset = load_csv_edge_dataset( # 你的其他参数(比如edge_path, node_path等) ) # 对数据集里的每个数据对象调整维度 for data in dataset: data.edge_label = data.edge_label.squeeze(dim=1)
补充说明
PyG的CSV加载工具读取单列标签时,默认会将每个标签包装成单独的一维数组,所以最终形成二维张量。而原始教程的数据集是把标签存储为一维张量,这就导致损失函数计算时输入(模型输出的一维张量)和目标(二维张量)维度不匹配,触发广播警告。调整后两者维度一致,就能解决这个问题。
内容的提问来源于stack exchange,提问作者wasif khan
相关产品推荐
相关产品推荐

