如何把供应链装进 PyG:图神经网络供应链优化实战指南
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
从哪个仓发货、备多少天的库存,这些决策都押在同一张网络上:谁供应谁、谁服务谁。这篇文章用 PyTorch Geometric(PyG)把这张网络装进一张图,训练一个图神经网络,让它输出运输成本与库存预测,全程只需要四段代码。
供应链为什么天然适合用一张图来表达
拿"确认一笔订单的发货方案"来说。客户下单后,你手上有三个候选仓。选哪个?看的不只是距离:这个仓现在有多少库存,上游供应商补货要几天,这条线路运费多少。
你会发现这些量分散在两类地方:库存、产能这类"状态"挂在节点上;距离、运费这类"代价"挂在节点之间的关系上。关系型数据库里,把它们拼齐要多张表 join;而图神经网络的消息传递,本来就是一次次"沿边向邻居要信息、再汇总到自己"。供应链的结构和这个过程是对得上的。
💡 记住一句话就够了:节点存状态,边存代价。
用 HeteroData 搭建供应链异构图
定义节点类型和 3 种边关系
PyG 的HeteroData相当于一个带隔间的图:每种节点、每种边各自拥有独立的特征和邻接信息,不用硬塞进同一张表。先定 3 种节点、3 种边(产品维度先合并进仓库特征,保持例子简单):
import torch from torch_geometric.data import HeteroData data = HeteroData() # 3 类节点,特征均为 16 维 data['supplier'].x = torch.randn(40, 16) # 供应商:产能、交期 data['warehouse'].x = torch.randn(12, 16) # 仓库:容量、当前库存 data['customer'].x = torch.randn(200, 16) # 客户:平均单量、区域 # 3 类边:edge_index 是 2 行张量,[源节点 id, 目标节点 id] data['supplier', 'supplies', 'warehouse'].edge_index = torch.randint(0, 40, (2, 60)) data['warehouse', 'serves', 'customer'].edge_index = torch.randint(0, 12, (2, 300)) data['customer', 'orders', 'warehouse'].edge_index = torch.randint(0, 200, (2, 300)) print(data.metadata()) # 所有节点/边类型的清单,后面要交给模型为什么不用单一的Data?因为供应商的 16 维和客户根本不是一个语义空间,强行共用一张邻接矩阵,模型就得靠特征自己猜"这条边是什么意思"。隔间结构把这个猜测成本直接删掉了。
训练一个能预测链路的图编码器
写一个可复用到所有边类型的两层 GraphSAGE
编码器只需要按"同构图"写一版。关键是SAGEConv((-1, -1), dim)里的-1:输入维度自动推断,你不用为每类节点硬编码。to_hetero拿到metadata()后,会按边类型把每一层复制成多份,每种关系各自训练一套参数:
from torch_geometric.nn import SAGEConv, to_hetero class GNNEncoder(torch.nn.Module): def __init__(self, dim): super().__init__() # (-1, -1):自动推断输入维度 self.conv1 = SAGEConv((-1, -1), dim) self.conv2 = SAGEConv((-1, -1), dim) def forward(self, x, edge_index): x = self.conv1(x, edge_index).relu() return self.conv2(x, edge_index) # 按边类型展开:每种关系拿到独立参数,aggr 指定邻居聚合方式 encoder = to_hetero(GNNEncoder(32), data.metadata(), aggr='sum')跑通第一次前向:输入data.x_dict和data.edge_index_dict,得到每类节点一个 32 维嵌入。到此为止,编码器并不知道你在做什么业务——它是通用的。
用一个读边两端嵌入的小 decoder 接业务
decoder 做的事很直白:取出边两端的嵌入,拼接,过两层 MLP,输出一个标量。标量是什么意思,由你放在边上的标签决定——可以是运费,也可以是补货天数:
class EdgeDecoder(torch.nn.Module): def __init__(self, dim): super().__init__() self.head = torch.nn.Sequential( torch.nn.Linear(2 * dim, dim), # 边两端嵌入拼接后线性变换 torch.nn.ReLU(), torch.nn.Linear(dim, 1), # 输出单个标量 ) def forward(self, z, src, dst): # z 为节点嵌入表,src/dst 是边的两端节点 id 序列 return self.head(torch.cat([z[src], z[dst]], dim=-1)).view(-1)训练就是常规的 MSE loss 加 Adam 循环。数据怎么切?官方示例 examples/hetero/hetero_link_pred.py 用T.RandomLinkSplit随机遮掉 20% 的边做验证和测试,推理时再放回去,值得照抄。
把预测接到真实业务量上
预测"仓库 → 客户"的运输成本
把每条线路的历史平均运费写进edge_label,用 MSE 训练上面这套 encoder + decoder。训练收敛后,模型对每条边都有一个预测值;对刚开的新线路、还没有历史数据的路径,它给出第一个估计值——链路预测在供应链里的价值就在这里:先估一条还没被实测过的边。
用 LinkNeighborLoader 处理会变动的网络
真实的网络不是静态的:线路会新增也会下线,仓库容量会调整。如果训练时用的是"全时间轴"的图,测试时模型其实见过未来。解决办法是让每条边带一个时间戳,采样时只取事件发生前的边。LinkNeighborLoader原生支持这件事:
from torch_geometric.loader import LinkNeighborLoader # 采样前先给边写时间戳:data['warehouse','serves','customer'].time = times loader = LinkNeighborLoader( data=data, num_neighbors=[8, 8], # 每跳采样 8 个邻居 batch_size=128, # 要预测的边类型 + 对应边集 edge_label_index=(('warehouse', 'serves', 'customer'), edge_index), edge_label_time=times - 1, # 取"事件之前"那一刻的图 time_attr='time', # 边上的时间属性名 temporal_strategy='last', # 同一邻居只保留最近一条边 )⚠️edge_label_time=times - 1这一行最容易被漏掉:不减 1,图就会在训练时"记得未来"。官方推荐系统示例 examples/hetero/recommender_system.py 用的就是同一套时序采样逻辑。
用指标检验模型值不值得上线
上线前先回答一个更尖锐的问题:模型比"永远预测历史平均"强多少?如果基线 RMSE 是 42 元,模型 43 元,架构再漂亮也别上。
- RMSE:
F.mse_loss(pred, target).sqrt(),和运费同单位,可以直接拿去和财务对账。 - MAE:平均绝对误差,个别极端线路不容易把它拉飞,适合看整体水位。
- Top-K:业务上你多半问的是"这单给前 3 个候选仓排序,最优选排进前十没有"。用
torch_geometric.metrics里的LinkPredPrecision/LinkPredRecall设k=10就能算。
规模方面:几万节点的网不必整图入训,LinkNeighborLoader每个 batch 只算两跳邻域;单机装不下时,torch_geometric/distributed/ 下的DistNeighborLoader会把图切到多台机器上,各自采样再汇总。
收个尾
整条链路就四步:HeteroData定义图 →to_hetero自动展开子网 → decoder 预测边标签 → 用 RMSE 和 Top-K 跟基线对账。想换"库存天数"或"准时率"做标签,改的只是edge_label那一行,架构不用动。你手上哪个决策环节最想交给一张图?评论区聊聊。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考