OSMNX转PyTorch Geometric报错:Could not infer dtype of Point及TGCN训练问题
解决OSMNX图转PyTorch Geometric格式及TGCN训练问题
报错原因
Could not infer dtype of Point 是因为你的NetworkX图里还保留着节点/边的geometry字段——也就是GeoPandas的Point/LineString类型对象,PyTorch Geometric(PyG)无法将这类非数值的几何类型转换成张量,必须先做清理或转换处理。
一、将OSMNX图转为PyG格式
1. 清理地理属性(必做)
先删除节点和边中的geometry字段,只保留你需要的数值特征:
import networkx as nx # 批量删除节点的geometry属性 nx.set_node_attributes(new_graph, {n: {} for n in new_graph.nodes()}, 'geometry') new_graph.remove_node_attributes('geometry') # 批量删除边的geometry属性(OSMNX生成的是多重图,需带keys=True) nx.set_edge_attributes(new_graph, {(u, v, k): {} for u, v, k in new_graph.edges(keys=True)}, 'geometry') new_graph.remove_edge_attributes('geometry')
2. 执行转换
清理完成后,再用from_networkx转换并指定要保留的特征:
import torch from torch_geometric.utils.convert import from_networkx pyg_graph = from_networkx(new_graph, group_node_attrs=["street_count"], group_edge_attrs=["length"]) print(pyg_graph)
可选:保留地理坐标作为节点特征
如果需要用到地理位置信息,可以把Point的x、y坐标提取为数值属性:
# 给节点添加坐标属性并删除原geometry for node in new_graph.nodes(data=True): if 'geometry' in node[1]: node[1]['x_coord'] = node[1]['geometry'].x node[1]['y_coord'] = node[1]['geometry'].y del node[1]['geometry'] # 转换时包含坐标特征 pyg_graph = from_networkx(new_graph, group_node_attrs=["street_count", "x_coord", "y_coord"], group_edge_attrs=["length"])
二、基于PyG开展TGCN训练
转成PyG格式后,可使用TGCNConv层构建模型,以下是一个基础的交通预测训练示例:
import torch.nn.functional as F from torch_geometric.nn import TGCNConv class TGCNModel(torch.nn.Module): def __init__(self, in_channels, hidden_channels, out_channels): super().__init__() self.tgcn1 = TGCNConv(in_channels, hidden_channels) self.tgcn2 = TGCNConv(hidden_channels, out_channels) def forward(self, x, edge_index): x = F.relu(self.tgcn1(x, edge_index)) x = self.tgcn2(x, edge_index) return x # 初始化模型(假设输入特征维度为1,比如单维度交通流量) model = TGCNModel(in_channels=1, hidden_channels=64, out_channels=1) optimizer = torch.optim.Adam(model.parameters(), lr=0.01) # 训练循环示例 for epoch in range(100): model.train() optimizer.zero_grad() # x为当前时间步的节点特征,shape: [节点数, 输入特征维度] out = model(x, pyg_graph.edge_index) # 计算MSE损失(根据你的任务调整损失函数) loss = F.mse_loss(out, y_true) loss.backward() optimizer.step() print(f'Epoch {epoch+1}, Loss: {loss.item():.4f}')
内容的提问来源于stack exchange,提问作者Wenyao Leo
相关产品推荐
相关产品推荐

