news 2026/9/8 17:59:15

3 步跑通:用 PyG 异构图给仓库到客户的运输成本算个明白账

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
3 步跑通:用 PyG 异构图给仓库到客户的运输成本算个明白账

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.metricsLinkPredPrecision/LinkPredRecall@K:Precision@10 回答「给每条线路推荐的 10 个候选合作方里平均几个真发生了往来」,Recall 回答「真实新合作被推荐列表覆盖了多少」。此时记得在 loader 里加neg_sampling=dict(mode='binary', amount=2)造负样本,让模型学会区分真合作和随机配对。

4 单机装不下时怎么办:分布式采样 + 脚本化上线

节点过百万、订单边动辄上亿时,整图进不了单机内存。PyG 的torch_geometric/distributed/给的是两级方案:

离线切图:Partitioner把节点、边、特征按分片落盘,每个partN/下是graph.ptnode_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),仅供参考

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

上地周边的硬科技创业社区:不是互联网玩法的那种

在北京海淀上地,大量创业载体集中涌现,很多创业者都会提出同一个问题:上地附近有什么硬科技创业社区吗?不是那种互联网运营的。互联网导向的创业社区普遍侧重流量运营、线上活动、新媒体曝光,更多服务消费互联网、软件…

作者头像 李华
网站建设 2026/9/8 17:52:14

剖析核心检查模块Check.js:启动白屏治理与前端架构设计

先说一个真实场景。某次线上启动白屏排查,业务侧转给我一份日志,里面只有一行:Check.js: Assertion Failed.。我盯着这行日志愣了几秒——Check.js是哪个文件?沿着仓库路径找进去,Source/Core/Check.js,安安…

作者头像 李华
网站建设 2026/9/8 17:51:08

嵌入式工程师必会:GPIO硬件结构与8种工作模式详解

做了十几年嵌入式开发,我面试过不少人,几乎每次都会从GPIO问起。GPIO这个外设入门时最容易点灯,往后越挖越深,越容易发现自己之前的理解只是半桶水。这篇文章把GPIO的硬件结构、8种工作模式、模式选择方法、实际调试思路一次讲透&…

作者头像 李华