news 2026/9/6 15:29:23

供应链里那些算不清的账,交给图神经网络:PyG 异构图运输成本预测实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
供应链里那些算不清的账,交给图神经网络:PyG 异构图运输成本预测实战

供应链里那些算不清的账,交给图神经网络:PyG 异构图运输成本预测实战

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

用 PyTorch Geometric(PyG,图神经网络库)把供应商、仓库、客户、产品装进同一张异构图,再预测「仓库→客户」这条边的运输成本:从建模、时序采样到分布式扩展和部署,本文走一遍完整链路。

速览:本文讲如何用 PyG 的HeteroData(异构图数据容器)+SAGEConv编码器做边级成本回归,并用LinkNeighborLoader处理随时间变化的运输关系。适合读者:有 Python/PyTorch 基础、想给自己的物流网络做预测或推荐的工程师。

1 把供应链装进一张图:HeteroData 异构图数据构建

传统做法是把供应链拆成一堆 Excel 和 SQL 视图分别分析,但「供应商产能不足 → 某仓缺货 → 改走另一条运输线 → 客户延期」这类影响是沿着关系传递的,表与表之间的传导在单表里根本看不到。图的价值就是把实体和它们的关系画在同一张纸上,让消息沿着关系传播。

用 PyG 的HeteroData建图,关键是节点和边都用「类型 + 索引」表达:

import torch from torch_geometric.data import HeteroData 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) # 体积重、温层、单价 # 边:2x2 的 index,第 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 顶上去。边类型不用穷举所有组合,只保留有业务语义的那几条。仓库自带示例 examples/hetero/hetero_link_pred.py 用的是「用户-评分-电影」异构图,把节点名换成供应链实体后结构完全通用。

2 边级回归:SAGEConv 编码器 + to_hetero 做运输成本预测

边级预测(预测某条线路的成本、时效、断供概率)是供应链里最常见的落地点。思路:先把每个节点编码成向量,再取边两个端点的向量拼起来过一个小 MLP,输出标量。

下面三段是模型核心,to_hetero负责把「同质 GNN」按元数据自动展开成异构图模型:

from torch_geometric.nn import SAGEConv, to_hetero from torch_geometric.transforms import RandomLinkSplit train_data, val_data, test_data = RandomLinkSplit( num_val=0.1, num_test=0.1, neg_sampling_ratio=0.0, # 回归任务不需要负样本 edge_types=[('warehouse', 'transports', 'customer')], rev_edge_types=[('customer', 'rev_transports', 'warehouse')], )(data) class GNNEncoder(torch.nn.Module): def __init__(self, hidden_channels, out_channels): super().__init__() self.conv1 = SAGEConv((-1, -1), hidden_channels) # -1: 输入维度自动推断 self.conv2 = SAGEConv((-1, -1), out_channels) def forward(self, x, edge_index): return self.conv2(self.conv1(x, edge_index).relu()) class EdgeDecoder(torch.nn.Module): def __init__(self, hidden_channels): super().__init__() self.lin1 = torch.nn.Linear(2 * hidden_channels, hidden_channels) self.lin2 = torch.nn.Linear(hidden_channels, 1) def forward(self, z_dict, edge_label_index): row, col = edge_label_index z = torch.cat([z_dict['warehouse'][row], z_dict['customer'][col]], dim=-1) return self.lin2(self.lin1(z).relu()).view(-1) class Model(torch.nn.Module): def __init__(self, hidden_channels): super().__init__() self.encoder = to_hetero(GNNEncoder(hidden_channels, hidden_channels), metadata=data.metadata(), aggr='sum') self.decoder = EdgeDecoder(hidden_channels) def forward(self, x_dict, edge_index_dict, edge_label_index): z_dict = self.encoder(x_dict, edge_index_dict) return self.decoder(z_dict, edge_label_index)

SAGEConv((-1, -1), ...)里的-1表示输入维度由数据推断,这样to_hetero才能自动给每种边类型配独立参数。注意RandomLinkSplitrev_edge_types要带上反向边,否则反向边会泄漏到训练集。

训练就是最普通的 MSE 回归;评估函数同时算了 RMSE 和 MAE,方便对到业务口径:

import torch.nn.functional as F model = Model(64) optimizer = torch.optim.Adam(model.parameters(), lr=0.01) def train(): model.train() optimizer.zero_grad() pred = model(train_data.x_dict, train_data.edge_index_dict, train_data['warehouse', 'transports', 'customer'].edge_label_index) loss = F.mse_loss(pred, train_data['warehouse', 'transports', 'customer'].edge_label) loss.backward() optimizer.step() return float(loss) @torch.no_grad() def test(data): model.eval() et = ('warehouse', 'transports', 'customer') pred = model(data.x_dict, data.edge_index_dict, data[et].edge_label_index) target = data[et].edge_label.float() return float(F.mse_loss(pred, target).sqrt()), \ float(F.l1_loss(pred, target)) for epoch in range(1, 201): print(train(), test(val_data))

训练完看test(test_data)。RMSE 和 MAE 都能直接换算成钱:测试集上 MAE=0.8 千元/单,乘以月单量 10 万就是「模型平均偏差 ≈ 80 万元/月」——这个数拿去和现在拍脑袋的固定报价比,就知道该不该用模型出价。判断好坏只看 test split,val 上的数字只用来早停。

3 动态运输网络:LinkNeighborLoader 时序邻居采样

静态切边有个坑:运输关系每天在变,「上周新开的一条线路」不该出现在训练里。仓库示例 examples/hetero/recommender_system.py 给出的做法是按时间戳切分:

from torch_geometric.loader import LinkNeighborLoader loader = LinkNeighborLoader( data=data, num_neighbors=[5, 5], edge_label_index=(('warehouse', 'transports', 'customer'), edge_index), edge_label_time=edge_time - 1, # 关键:-1 防止采样到未来边 time_attr='time', temporal_strategy='last', # 每跳只取截断时刻之前的邻居 batch_size=256, shuffle=True, )

temporal_strategy='last'让每一跳采样都受edge_label_time约束:只采到「预测时点」之前的历史边,从机制上杜绝未来信息泄漏——这在物流场景里比模型本身更容易出错。

评估侧换成推荐口径:LinkPredPrecision(k)/LinkPredRecall(k)(来自torch_geometric.metrics)。Precision@10 可以解读为「给每条线路推荐 10 个候选合作方,平均有 K 个是真发生了往来的」;召回率则回答「真实合作被推荐列表覆盖了多少」。链路预测任务记得在 loader 里加neg_sampling=dict(mode='binary', amount=2)造负样本。

4 图大到装不下内存:分布式邻居采样

当节点过百万,单机装不下整张图,torch_geometric/distributed/提供两级扩展:

  1. 离线切图Partitioner把节点和特征按分片落盘(每个partN/下是graph.ptnode_feats.pt);
  2. 在线采样DistNeighborLoader绑定本分片,本地邻居直接读,跨分片邻居走 RPC 从远端拉。

效果是采样开销从「全图」降到「本机分片 + 一跳远程」,训练吞吐随机器数近似线性扩展。对供应链这类边数远超单机的网络(大客户订单边动辄上亿),这一步基本是必选项。

5 上线部署:torch.jit 脚本化导出与加载

训练完的模型用torch.jit.script导出,推理侧不再依赖 Python 训练环境(参考 examples/jit/gin.py 的做法):

scripted = torch.jit.script(model) torch.jit.save(scripted, 'supply_chain_model.pt') loaded = torch.jit.load('supply_chain_model.pt') pred = loaded(x_dict, edge_index_dict, edge_label_index)

注意导出的是「编码器 + 解码器」整体,输入仍然是x_dict/edge_index_dict,线上服务把特征拼装好直接喂入即可;如果线上只更新编码器(特征变了但解码关系不变),也可以只导编码器单独服务。

6 要点回顾与下一步

要点回顾:

  1. 建模HeteroData表达多节点/多边类型;边用 2×E 的 index,节点特征做归一化
  2. 模型SAGEConv((-1, -1), ...)让维度自动推断,to_hetero按元数据展开成异构图
  3. 边级回归:拼接两端点向量 → 小 MLP 输出标量,MSE 训练,RMSE/MAE 对账到业务金额
  4. 时序采样LinkNeighborLoader+temporal_strategy='last'+edge_label_time - 1防未来泄漏
  5. 扩展torch_geometric/distributed/做切图与跨机采样;torch.jit脚本化部署

下一步可以做的事(按优先级):

  • 给解码器加多任务头:同一个z_dict上同时预测成本、时效、断供概率,用不同的 decoder 共享编码器
  • 链路预测任务加neg_sampling调参,负样本比例对 Precision@K 影响很大
  • 把 MAE 换成业务可解释的损失(如分段线性),让模型在「大客户线路」上偏差更小

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/6 15:28:19

CentOS 7.9 停止维护(2024-6-30)后可用在线yum源 —— 筑梦之路

众所周知,centos 7 在2024年6月30日,生命周期结束,官方不再进行支持维护,而很多环境一时之间无法完全更新替换操作系统,因此对于yum源还是需要的,特别是对于互联网环境来说,在线yum源使用方便很…

作者头像 李华
网站建设 2026/9/6 15:27:06

5分钟看懂Coze Studio监控面板:四大核心指标完整指南

5分钟看懂Coze Studio监控面板:四大核心指标完整指南 【免费下载链接】coze-studio An AI agent development platform with all-in-one visual tools, simplifying agent creation, debugging, and deployment like never before. Coze your way to AI Agent creat…

作者头像 李华
网站建设 2026/9/6 15:26:42

docker操作文档

一、升级数据库 A.先网管备份下sql B.然后docker操作 1.移除mysql的docker容器 docker rm -f mars-mysql-server2.将docker_data里面的mysql删除或者改名称 mv mysql mysql013.进去docker_compose重新编译docker docker-compose up -d --buildC.最后使用备份的sql文件恢复数据库…

作者头像 李华
网站建设 2026/9/6 15:24:45

Buzz 离线语音转文字实操指南:5 分钟转好、改好、导出字幕

Buzz 离线语音转文字实操指南:5 分钟转好、改好、导出字幕 【免费下载链接】buzz Buzz transcribes and translates audio offline on your personal computer. Powered by OpenAIs Whisper. 项目地址: https://gitcode.com/GitHub_Trending/buz/buzz Buzz 是…

作者头像 李华
网站建设 2026/9/6 15:21:34

Multisim仿真对比:二极管包络检波与同步检波原理及失真分析

简介:二极管包络检波与同步检波仿真实验报告,面向通信电子线路课程学习者及高频电路实验人员,完整呈现调幅波解调的仿真验证过程。报告基于 EWB 软件搭建二极管包络检波器与双边带调幅同步检波电路,包含调幅度 0.8 下的正常波形、…

作者头像 李华