Pytorch Geometric中如何向from_networkx传递第二、第三位置参数?
PyTorch Geometric from_networkx 传参报错解决方案
报错根因
torch_geometric.utils.from_networkx()仅第一个参数(传入networkx图对象)支持位置传参,其余所有参数均为关键字参数,不能通过位置顺序传入。你直接将属性列表作为第二个位置参数传入,就会触发参数数量不匹配的错误。
正确调用方式
如果要提取指定节点属性,将属性列表传给group_node_attrs关键字参数即可:
# 提取spin节点属性并整合到Data的x属性中 data = from_networkx(I, group_node_attrs=["spin"]) print(data)
参数说明
- 当
group_node_attrs传入属性名列表时,对应节点属性会被拼接后存入生成Data对象的x字段 - 如果需要保留属性名作为
Data对象的独立字段,设置group_node_attrs=False即可,此时可以直接通过data.spin访问对应属性
替代解决方案(适配旧版本PyG)
如果你的PyG版本过低不支持上述参数,可以手动构造PyG的Data对象:
import torch from torch_geometric.data import Data # 生成边索引 edge_index = torch.tensor(list(I.edges), dtype=torch.long).t().contiguous() # 按节点顺序提取spin属性 node_spin = torch.tensor([I.nodes[node]["spin"] for node in I.nodes], dtype=torch.float).reshape(-1, 1) # 构造Data对象 data = Data(x=node_spin, edge_index=edge_index)
内容的提问来源于stack exchange,提问作者CR-97
相关产品推荐
相关产品推荐

