PyTorch Geometric 异构图实战完整指南:3 段代码预测每条运输线路的运费
发布时间:2026/9/8 16:20:03 作者:尧图编辑部 阅读量:1,286

PyTorch Geometric 异构图实战完整指南3 段代码预测每条运输线路的运费【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric如果你管过物流网络大概被这样的报价问题折磨过一条「仓库 → 客户」线路该报多少钱它不只取决于里程还受供应商产能、仓库库存、货物属性牵制——这些变量散在不同表里传统报表工具只能各算各的。本文用图神经网络库 PyTorch GeometricPyG把供应商、仓库、客户、产品装进同一张异构图对每条运输边做成本回归预测全文给出 3 段核心代码建图、建模训练、时序采样外加分布式扩展与脚本化部署的做法。 用HeteroData把多张业务表拼成一张「多角色」图 用SAGEConvto_hetero写出边级运费回归模型LinkNeighborLoader按时钟采样邻居杜绝未来信息泄漏1 运费为什么是一张「多角色」图的事先说痛点供应商缺产能 → 某仓断货 → 改走备用线路 → 客户侧延期。这条因果链跨了三张表在单表 SQL 里根本串不起来。图模型的思路是把所有实体画在同一张图上让信息顺着边一跳一跳地扩散节点最终「知道」了远端发生的事情。异构图HeteroData可以类比为「一张贴了岗位标签的流程图」节点分角色供应商、仓库、客户、产品边也分种类供货、存货、运输。同一套 GNN 代码不需要为每种角色单独写这是它相对纯关系模型最大的省心之处。2 把业务表拼成 HeteroData这一节解决「数据从哪来、怎么装」。建图就一个容器对象节点和边都用「类型名 2 行索引」表达data HeteroData() data[supplier].x torch.randn(120, 8) # 产能、区位、履约率 data[warehouse].x torch.randn(30, 8) # 库容、周转天数、租金 data[customer].x torch.randn(5000, 8) # 下单频次、账期、区域 data[product].x torch.randn(300, 8) # 体积重、温层、单价 # 每条边是 2xN 张表第 0 行是起点节点第 1 行是终点节点 data[supplier, supplies, warehouse].edge_index sup_wh_idx data[warehouse, stores, product].edge_index wh_prod_idx data[warehouse, transports, customer].edge_index wh_cust_idx特征直接从 ERP/WMS 取现成字段做z-score归一化即可。两类细节容易踩坑边类型别贪多只保留有业务语义的关系供货、存货、运输把所有两两组合都建成边只会引入噪声。缺特征的节点纯关系型节点先用独热 ID 顶上去保证x维度齐全。仓库里 examples/hetero/hetero_link_pred.py 是「用户-评分-电影」的异构图边预测示例把角色名换成上面四个供应链实体结构完全通用。3 边级回归同一段 GNN 代码自动展开成异构图模型这一节解决「模型怎么写」。目标是为每条「仓库 → 客户」边输出一个标量运费。套路是两段式编码器把节点压成向量解码器取边两端向量拼接后过一个小 MLP 出数。切分放在建模之前注意rev_edge_types必须登记反向边否则反向边会漏进训练集from torch_geometric.transforms import RandomLinkSplit train_data, val_data, test_data RandomLinkSplit( num_val0.1, num_test0.1, neg_sampling_ratio0.0, # 回归任务不采负样本 edge_types[(warehouse, transports, customer)], rev_edge_types[(customer, rev_transports, warehouse)], )(data)模型本体只写一份「同质」代码to_hetero会按data.metadata()自动给每种节点/边类型复制独立参数from torch_geometric.nn import SAGEConv, to_hetero class GNNEncoder(torch.nn.Module): def __init__(self, hidden, out): super().__init__() # -1 表示维度从数据推断to_hetero 才能自动展开 self.conv1 SAGEConv((-1, -1), hidden) self.conv2 SAGEConv((-1, -1), out) def forward(self, x, edge_index): return self.conv2(self.conv1(x, edge_index).relu()) class Model(torch.nn.Module): def __init__(self, hidden): super().__init__() self.encoder to_hetero(GNNEncoder(hidden, hidden), data.metadata(), aggrsum) self.lin1 torch.nn.Linear(2 * hidden, hidden) self.lin2 torch.nn.Linear(hidden, 1) def forward(self, x_dict, edge_index_dict, edge_label_index): z self.encoder(x_dict, edge_index_dict) row, col edge_label_index z torch.cat([z[warehouse][row], z[customer][col]], dim-1) return self.lin2(self.lin1(z).relu()).view(-1)训练就是常规 MSE 回归optimizer.zero_grad()→ 前向 →F.mse_loss(pred, target).backward()→optimizer.step()用Adam、学习率0.01跑 200 个 epoch 即可收敛。注意解码器拼的是z[warehouse][row]和z[customer][col]——边的起点向量取 row、终点向量取 col别写反。4 把误差换算成每月多少钱这一节解决「数字好不好怎么判断」。评估同时看两个口径F.mse_loss(pred, target).sqrt()得RMSE偏差大的线路会被放大反映「最离谱的单子偏多少」F.l1_loss(pred, target)得MAE单笔平均偏差最容易对账。换算到业务侧很直接假设测试集 MAE 0.8 千元/单月单量 10 万那么模型的系统性平均偏差约为 0.8 × 10 万 80 万元/月。拿这个数和现有拍脑袋的固定报价方案比模型该不该上线就有量化答案了。判断模型好坏只认 test splitval split 的数字留给早停策略不写进汇报。5 上生产前的三道关5.1 线路天天在变按时钟做邻居采样运输关系是动态的「上周刚开的新线路」不能出现在「上周」的训练数据里。做法是LinkNeighborLoader 时间属性仓库的 examples/hetero/recommender_system.py 就是这个路子from torch_geometric.loader import LinkNeighborLoader loader LinkNeighborLoader( data, num_neighbors[5, 5], edge_label_index((warehouse, transports, customer), edge_index), edge_label_timeedge_time - 1, # 关键减 1采样只看到「之前」的边 time_attrtime, temporal_strategylast, # 每跳只取截断时刻之前的邻居 batch_size256, shuffleTrue, )temporal_strategylast让每一跳采样都被edge_label_time截断未来信息在机制上进不来——这类泄漏在动态网络里比模型本身更容易出错。如果任务从回归换成链路预测预测「会不会新开这条线」评估换成torch_geometric.metrics里的LinkPredPrecision(k)/LinkPredRecall(k)并给 loader 加neg_samplingdict(modebinary, amount2)造负样本负样本比例对 PrecisionK 影响很大值得调参。5.2 图大到装不进内存切图 跨机采样节点过百万后单机放不下torch_geometric/distributed/提供两级方案离线切图Partitioner把节点和特征分片落盘每个partN/目录下是graph.pt边索引与node_feats.pt节点特征在线采样DistNeighborLoader绑定本机分片本地邻居直接读跨分片邻居走 RPC 从远端机器拉取。采样开销从「扫全图」降到「本机分片 一跳远程」吞吐随机器数近似线性扩展。订单边动辄上亿的供应链网络这一步基本是必选项。5.3 脚本化导出推理侧甩掉训练环境torch.jit.script(model)torch.jit.save()导出后线上服务用torch.jit.load()加载即可不再依赖 Python 训练栈参考 examples/jit/gin.py。导出的是编码器 解码器整体输入仍是x_dict/edge_index_dict/edge_label_index线上把特征拼装好直接喂入。如果只有特征在变、边关系不变也可以只导出编码器单独提供服务。延伸方向多任务头同一个z_dict上挂多个 decoder共享编码器同时预测成本、时效、断供概率链路预测的neg_sampling比例调参观察 PrecisionK 变化曲线把 MSE 换成分段线性损失压低大客户线路上的偏差用 examples/multi_gpu/ 里的distributed_sampling.py跑通多机采样全流程【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考