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

如何将CSV数据集转换为PyG支持的标准图数据格式

CSV转PyG图数据实现方案

边索引(edge_index)正确构建

PyG 要求输入的edge_index是形状为[2, 总边数]的长整型张量,第一维度对应边的源节点、第二维度对应边的目标节点。
你的edge_pairs字段是逐节点存储的连边,处理时注意两个常见问题:

  • CSV读入时该字段是字符串格式,需要先解析成原生Python列表
  • 无向图场景下同一条边会在两个关联节点的行中各存一次,必须去重,否则边数会虚高一倍
    处理逻辑:
  1. 遍历所有行的edge_pairs字段,用ast.literal_eval把字符串转成嵌套列表
  2. 逐边收集(u, v)对,无向图同步加入反向边(v, u),用集合做去重
  3. 把去重后的边列表转置成[2, num_edges]格式,转成torch.long类型张量即可
    如果你的节点编号不是从0开始连续排列,需要先做一次重映射,把所有节点id映射到0 ~ 节点总数-1的连续区间,否则PyG会触发索引越界错误。

节点特征(x)选取方案

x要求是形状为[节点总数, 特征维度]的浮点型张量,你当前数据集没有现成特征列,根据任务选常用方案即可:

  • 如果s_key/identifier是节点的固有属性(比如节点类别、所属分组):直接用标签编码/独热编码转成张量作为特征,字符串类型属性先做编码再输入
  • 如果没有任何节点属性、做图表示学习/节点分类任务:可以用单位矩阵(每个节点对应一个one-hot向量,维度等于节点总数)作为初始特征,也可以提取节点度、PageRank值、聚类系数这类拓扑统计量作为特征
  • 不要用index列直接当特征,这一列只是节点编号,没有实际语义信息。

节点标签(y)选取方案

y完全和你的下游任务绑定,没有固定取值规则:

  • 做节点分类任务:用你提前标注好的节点类别标签,多分类场景下形状为[节点总数]的长整型张量,多标签场景下为[节点总数, 类别数]的浮点张量
  • 做链路预测任务:不需要全局y,后续数据集划分时单独构造正负样本边即可
  • 做图预训练、拓扑结构分析任务:暂时不需要y可以留空,或者把你需要预测的节点属性(比如identifier对应的类别)赋值给y即可。

完整可运行代码

import pandas as pd
import torch
from torch_geometric.data import Data
import ast
from sklearn.preprocessing import LabelEncoder

# 读取CSV数据
df = pd.read_csv("your_graph_dataset.csv")
num_nodes = len(df)  # 前提是index列从0开始连续无断号

# 构建edge_index
edge_set = set()
for edge_col_val in df["edge_pairs"]:
    # 解析字符串格式的边列表
    edge_pairs = ast.literal_eval(edge_col_val)
    for u, v in edge_pairs:
        edge_set.add((u, v))
        # 有向图请删除下面这行,不要加反向边
        edge_set.add((v, u))
# 转成PyG要求的edge_index格式
edge_index = torch.tensor(list(edge_set), dtype=torch.long).t().contiguous()

# 构建节点特征x,以下两种方案二选一
# 方案1:用s_key属性编码作为特征
s_key_encoder = LabelEncoder()
s_key_feat = s_key_encoder.fit_transform(df["s_key"]).reshape(-1, 1)
x = torch.tensor(s_key_feat, dtype=torch.float)
# 方案2:无属性场景用one-hot初始特征
# x = torch.eye(num_nodes, dtype=torch.float)

# 构建节点标签y,替换成你实际的标签列即可
id_encoder = LabelEncoder()
y = torch.tensor(id_encoder.fit_transform(df["identifier"]), dtype=torch.long)

# 组装成PyG标准Data对象
data = Data(x=x, edge_index=edge_index, y=y)
print(data) # 打印验证数据格式是否正确

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 01:33:20