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

PyTorch Geometric中MNISTSuperpixels节点特征修改无效求助

解决PyTorch Geometric中MNISTSuperpixels节点特征替换无效的问题

问题背景

尝试替换PyTorch Geometric(PyG)中MNISTSuperpixels数据集的节点特征:先用CNN提取MNIST图像的10维特征并保存到文件,再将这些特征替换原数据集的节点特征,但修改后dataset[0].x仍保留原特征,num_features也未更新。

问题原因

  1. 迭代副本而非原数据对象:遍历zip(dataset, features)时,拿到的data是数据集元素的临时副本,修改副本不会同步到原数据集。
  2. 特征类型与维度不匹配:直接赋值Python列表给data.x不符合PyGData对象的要求(需为PyTorch张量),且提取的图像级10维特征未适配节点数量的维度(原data.x为[num_nodes, 1],需转为[num_nodes, 10])。
  3. 未更新数据集全局特征数:dataset.num_features是初始化时设置的属性,修改单个data.x不会自动更新该值。

修正后的代码

import torch
from torch_geometric.datasets import MNISTSuperpixels

# 加载MNIST Superpixels数据集
dataset = MNISTSuperpixels(root='.', train=True)

# 从文件加载提取的特征并转为张量
features = []
with open('features.txt', 'r') as f:
    for line in f:
        parts = line.strip().split(',')
        label = int(parts[0])
        # 将特征转为float类型张量
        img_feature = torch.tensor([float(x) for x in parts[1:]], dtype=torch.float)
        features.append((label, img_feature))

# 验证特征数量与数据集样本数一致
assert len(features) == len(dataset), "特征数量与数据集样本数不匹配"

# 逐个替换样本的节点特征
for idx in range(len(dataset)):
    data = dataset[idx]
    _, img_feature = features[idx]
    # 将图像级10维特征扩展为每个节点的特征:[10] → [num_nodes, 10]
    data.x = img_feature.repeat(data.num_nodes, 1)

# 更新数据集的全局特征数属性
dataset.num_features = 10

# 验证修改结果
print("修改后的节点特征形状:", dataset[0].x.shape)
print("修改后的数据集特征数:", dataset.num_features)

关键修改说明

  • 直接操作原数据对象:通过dataset[idx]索引访问数据集元素,修改后直接生效(MNISTSuperpixels是InMemoryDataset,存储的是Data对象列表,索引访问为直接引用)。
  • 特征类型转换:将提取的特征转为torch.Tensor,符合PyGData.x的类型要求。
  • 维度适配:用repeat(data.num_nodes, 1)将单张图像的10维特征复制到每个节点,确保与原节点特征的维度结构匹配(每个节点对应10维特征)。
  • 更新全局特征数:手动修改dataset.num_features,保证数据集的元信息与实际特征维度一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 18:50:24