供应链里那些算不清的账,交给图神经网络: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才能自动给每种边类型配独立参数。注意RandomLinkSplit时rev_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/提供两级扩展:
- 离线切图:
Partitioner把节点和特征按分片落盘(每个partN/下是graph.pt与node_feats.pt); - 在线采样:
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 要点回顾与下一步
要点回顾:
- 建模:
HeteroData表达多节点/多边类型;边用 2×E 的 index,节点特征做归一化 - 模型:
SAGEConv((-1, -1), ...)让维度自动推断,to_hetero按元数据展开成异构图 - 边级回归:拼接两端点向量 → 小 MLP 输出标量,MSE 训练,RMSE/MAE 对账到业务金额
- 时序采样:
LinkNeighborLoader+temporal_strategy='last'+edge_label_time - 1防未来泄漏 - 扩展:
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),仅供参考