3 步跑通:用 PyG 异构图给仓库到客户的运输成本算个明白账
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
这篇实战带你用 PyTorch Geometric(PyG,图神经网络库)把供应商、仓库、客户、产品建成一张异构图,对「仓库→客户」线路做供应链运输成本预测:从建模、防泄漏采样到分布式上线的完整链路,读完你能直接套到自己的物流网络上。
1 异构图数据怎么建:先想清楚单表为什么不够
把供应商、仓库、客户、产品拆成几张 Excel 分别分析,跨实体的传导就断了——供应商产能不足拖垮某个仓,仓改走另一条运输线,最终落到客户延期,这条链在单表里根本看不见。图神经网络的价值就是让信息沿着关系走。
你可以把HeteroData理解成一张「分线路的地铁图」:节点和边都按类型分开存,每条边用 2×E 的索引记录起点和终点。
data = HeteroData() data['supplier'].x = zscore(供应商特征) # 产能、区位、履约率 data['warehouse'].x = zscore(仓库特征) # 库容、周转天数、租金 data['customer'].x = zscore(客户特征) # 下单频次、账期、区域 data['supplier', 'supplies', 'warehouse'].edge_index = sup_wh_idx data['warehouse', 'transports', 'customer'].edge_index = wh_cust_idx节点特征从 ERP/WMS 取现成字段做 z-score 归一化;个别纯关系型节点没有特征列,用独热 ID 顶上去即可,消息传递照样能跑。边类型不用穷举所有组合,只留业务说得通的几条:谁供谁、谁运给谁,够了。仓库里 examples/hetero/hetero_link_pred.py 用「用户-评分-电影」演示了完全同构的写法,把实体名一换就是你的供应链。
2 边级回归模型怎么写:误差怎么对账成钱
预测目标是最常见的落地点——线路成本回归:编码器把节点压成向量,解码器取边两端点的向量拼起来,过一个小 MLP 输出一个标量。
关键在to_hetero:你只管写一个同质 GNN,它按data.metadata()自动展开成「每种边类型一套参数」的异构模型。为此SAGEConv的输入维度要写-1,表示从数据推断,这样展开时才不会对不上尺寸。
class Encoder(torch.nn.Module): def __init__(self, h): super().__init__() self.conv1 = SAGEConv((-1, -1), h) self.conv2 = SAGEConv((-1, -1), h) def forward(self, x, edge_index): return self.conv2(self.conv1(x, edge_index).relu()) model = torch.nn.Module() model.encoder = to_hetero(Encoder(64), data.metadata(), aggr='sum') model.decoder = EdgeDecoder(64) # 拼接两端点向量 + 两层 MLP,输出 1 维训练用 MSE 就行,真正值钱的是评估口径。RMSE/MAE 算出来后直接对账:测试集 MAE 若是每单 0.8 千元,月单量 10 万,模型平均偏差约 80 万元/月——拿这个数去对比现在拍脑袋固定报价的偏差,模型省不下这个数就别上线。
⚠️ 判断模型好坏只看 test split;val 的数字只用来早停,别拿它汇报。
数据切分用RandomLinkSplit:回归任务没有正负样本之分,neg_sampling_ratio=0;但rev_edge_types必须显式声明反向边类型,让切分时同步处理,这是后面要专门讲的坑。
3 反向边泄漏怎么防、时序泄漏怎么断
坑 1:反向边泄漏。为了消息传递顺畅,图通常补了反向边(如「客户→仓库」)。切分时只声明正向边,反向边就原封不动留在训练图里,测试边在结构上提前曝光。
一行修复:切分时把反向边类型一并交给RandomLinkSplit——
train_data, val_data, test_data = T.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)坑 2:时序泄漏。运输关系每天在变,上周才开的新线路不该出现在「预测这周」的训练里,否则模型背的是答案。
一行修复:LinkNeighborLoader里把edge_label_time减 1,配合temporal_strategy='last',每跳采样只取截断时刻之前的邻居——参考 examples/hetero/recommender_system.py 的做法:
loader = LinkNeighborLoader( data=data, num_neighbors=[5, 5], edge_label_index=(('warehouse', 'transports', 'customer'), edge_index), edge_label_time=edge_time - 1, # 只采「预测时点」之前的历史边 time_attr='time', temporal_strategy='last', batch_size=256, )如果任务从回归换成链路预测(预测哪对仓库-客户会新增往来),评估口径换成torch_geometric.metrics的LinkPredPrecision/LinkPredRecall@K:Precision@10 回答「给每条线路推荐的 10 个候选合作方里平均几个真发生了往来」,Recall 回答「真实新合作被推荐列表覆盖了多少」。此时记得在 loader 里加neg_sampling=dict(mode='binary', amount=2)造负样本,让模型学会区分真合作和随机配对。
4 单机装不下时怎么办:分布式采样 + 脚本化上线
节点过百万、订单边动辄上亿时,整图进不了单机内存。PyG 的torch_geometric/distributed/给的是两级方案:
离线切图:Partitioner把节点、边、特征按分片落盘,每个partN/下是graph.pt和node_feats.pt:
在线采样:DistNeighborLoader绑定本分片,本地邻居直读,跨分片邻居走 RPC 向远端机器拉一跳。
效果是采样开销从「全图」压到「本机分片 + 一跳远程」,吞吐随机器数近似线性涨——对供应链这种边数远超单机的网络,这步基本是必选项。
部署走脚本化:训练完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,线上服务把特征拼好直接喂入。若线上只是特征更新了(解码关系没变),也可以只导编码器单独服务,更新成本更低。
5 三个问题收尾:往哪扩展、损失怎么改、负样本比怎么定
Q:想同时预测成本、时效、断供概率,要训练三个模型吗?
不用。同一套z_dict上挂多个小解码头即可,编码器共享——图结构编码最贵,三个头分摊一份编码,而且任务间梯度互相约束,节点表征比单任务更稳。
Q:MSE 对大客户线路太「一视同仁」了怎么办?
把损失换成业务可解释的形态,比如分段线性或按金额加权:高客单线路每错一千元比小线路代价大得多,权重跟着金额走,模型会主动牺牲小单精度去压低大线路的偏差。
Q:负样本比例对 Precision@K 影响大吗?
很大。neg_sampling=dict(mode='binary', amount=2)是每条正样本配两条负样本;比例决定模型学到的「随机配对有多常见」,从而移动判定阈值。业务里真实合作稀疏,负样本给少了 Precision 虚高,给多了召回掉——用验证集扫一遍,看 Precision/Recall 的平衡点定。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考