news 2026/9/8 16:18:00

PyTorch Geometric 异构图实战:把供应链运输成本预测完整落地指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch Geometric 异构图实战:把供应链运输成本预测完整落地指南

PyTorch Geometric 异构图实战:把供应链运输成本预测完整落地指南

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

速览:用 PyTorch Geometric(PyG)把供应商、仓库、客户、产品四类实体装进同一张异构图,对「仓库→客户」边做运输成本回归——从最小可跑版本、切分可信度、时序采样防泄漏,一路走到分布式采样与 torch.jit 交付,链路一次走全。适合有 Python/PyTorch 基础、没碰过图学习的工程师,读完可以直接照抄骨架去接自己的物流数据。

一、先跑通:PyG 异构图最小可跑版本

不等概念铺完,先让代码转起来。PyG 的HeteroData用「节点类型 + 边类型」组织数据,边统一是 2×E 的索引:第 0 行起点、第 1 行终点。

import torch from torch_geometric.data import HeteroData data = HeteroData() data['supplier'].x = torch.randn(120, 8) # 产能、区位、履约率(z-score) data['warehouse'].x = torch.randn(30, 8) # 库容、周转天数、租金 data['customer'].x = torch.randn(5000, 8) # 下单频次、账期、区域 data['product'].x = torch.randn(300, 8) # 体积重、温层、单价 data['supplier', 'supplies', 'warehouse'].edge_index = sup_wh # 2xN data['warehouse', 'stores', 'product'].edge_index = wh_prod # 2xN data['warehouse', 'transports', 'customer'].edge_index = wh_cust # 待预测边

模型侧只需两段:SAGEConv编码器负责消息传递,解码器把边两端点向量拼起来过 MLP 输出标量。

from torch_geometric.nn import SAGEConv, to_hetero class Encoder(torch.nn.Module): def __init__(self, h, o): super().__init__() self.conv1 = SAGEConv((-1, -1), h) # -1 表示维度由数据推断 self.conv2 = SAGEConv((-1, -1), o) def forward(self, x, edge_index): return self.conv2(self.conv1(x, edge_index).relu()) class EdgeDecoder(torch.nn.Module): def __init__(self, h): super().__init__() self.mlp = torch.nn.Sequential( torch.nn.Linear(2 * h, h), torch.nn.ReLU(), torch.nn.Linear(h, 1)) def forward(self, z, edge_index): row, col = edge_index z = torch.cat([z['warehouse'][row], z['customer'][col]], -1) return self.mlp(z).view(-1) # 标量:单均运输成本 class CostModel(torch.nn.Module): def __init__(self, h): super().__init__() self.enc = to_hetero(Encoder(h, h), data.metadata(), aggr='sum') self.dec = EdgeDecoder(h) def forward(self, x_dict, e_dict, e_label): return self.dec(self.enc(x_dict, e_dict), e_label)
import torch.nn.functional as F model = CostModel(64) opt = torch.optim.Adam(model.parameters(), lr=0.01) et = ('warehouse', 'transports', 'customer') pred = model(data.x_dict, data.edge_index_dict, data[et].edge_label_index) loss = F.mse_loss(pred, data[et].edge_label) # 回归,MSE 起步 loss.backward() opt.step()

跑通之后再看两个「为什么」。其一,成本不是单边属性:供应商产能掉、某仓缺货、改走另一条线,影响是沿着边一跳跳传过去的,任何单表 SQL 都截不住这条链;图的价值就是让消息沿关系走。其二,边级预测的本质是「读两端猜整条边」——编码器把节点压成向量,解码器只消费拼接后的端点向量,所以模型规模可以和边数解耦。

二、把预测做可信:反向边防泄漏与指标口径

模型能跑不等于结果能用。先看最隐蔽的坑:训练/验证/测试要按「边」切,而不是按「节点」。

切边只切正向边、反向边原封不动留在原图,采样子图会把测试边带进训练批次——测试 RMSE 虚低,上线即失真。

RandomLinkSplitrev_edge_types把反向边一起切走:

from torch_geometric.transforms import RandomLinkSplit train, val, test = RandomLinkSplit( num_val=0.1, num_test=0.1, neg_sampling_ratio=0.0, # 回归不需要负样本 edge_types=[et], # 待预测边 rev_edge_types=[('customer', 'rev_transports', 'warehouse')], )(data)

评估侧同时算 RMSE 和 MAE,并且直接换算成钱,方便和业务对账:

@torch.no_grad() def evaluate(d): model.eval() p = model(d.x_dict, d.edge_index_dict, d[et].edge_label_index) y = d[et].edge_label.float() rmse = float(F.mse_loss(p, y).sqrt()) # 千元/单 mae = float(F.l1_loss(p, y)) return rmse, mae

假设测试集 MAE 是 0.4 千元/单,月均单量 8 万,平均偏差约为 3.2 万元/月。拿这个数和现行固定报价方案的实际偏差比,值不值得上线一眼可见。

val 的数字只用于早停,对外汇报一律取 test split;拿 val 当成绩,等于用考卷原题备考。

best, bad = float('inf'), 0 for epoch in range(1, 201): train_one_epoch(train) # 省略:前向 + MSE + step rmse_v, _ = evaluate(val) best, bad = (min(best, rmse_v), bad + 1) if rmse_v < best else (best, bad + 1) if bad >= 8: break # 早停,8 轮无改善 print(evaluate(test)) # 只信这里

三、时间进场:时序邻居采样如何防未来泄漏

运输关系每天都在变,上周才开的线路不该出现在训练样本里。静态切边到这里就到头了,需要换LinkNeighborLoader按时序采样。

from torch_geometric.loader import LinkNeighborLoader loader = LinkNeighborLoader( data, num_neighbors=[5, 5], edge_label_index=(et, wh_cust), # 待预测边 edge_label_time=edge_time - 1, # -1:只允许采样到过去 time_attr='time', temporal_strategy='last', # 每跳截断在预测时点之前 batch_size=256, shuffle=True, )

temporal_strategy='last'配合edge_label_time意味着:每一跳采样都只能拿到该边「预测时点」之前已经存在的关系,未来边在机制上不可见,而不是靠事后清洗。这在物流数据里比模型结构更容易翻车——线上口径对不上,多半先查这里。

评估口径换成推荐式的:torch_geometric.metrics提供LinkPredPrecision(k)LinkPredRecall(k)。Precision@20 读作「给每条线路推 20 个候选合作方,平均几个是真发生过往来的」;召回回答「真实合作被 Top-20 覆盖了多少」。转成链路预测任务时记得在 loader 里加neg_sampling=dict(mode='binary', amount=2)造负样本——负样本比例对 Precision@k 影响很大,要调参而不是写死。

四、上规模与交付:分布式采样和 torch.jit 导出

节点和边超出单机内存之后,PyG 的torch_geometric/distributed/分两级扩展。

先用Partitioner离线切图,产物是META.json、节点/边映射表加若干分片,每片带graph.ptnode_feats.pt

from torch_geometric.distributed import Partitioner Partitioner(data, num_parts=4, root='./parts').partition() # 落盘:META.json + node_map/ + edge_map/ + part0..part3/

训练侧换成DistNeighborLoader,绑定本分片后,本地邻居直接读盘,跨分片邻居走 RPC 异步拉取,采样范围从全图收窄到「本机分片 + 一跳远程」,吞吐随机器数近似线性扩——订单边动辄上亿时这一步基本是必选项。

交付侧,模型用torch.jit.script整体导出,推理环境不再依赖训练时的依赖树:

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

导出的对象是「编码器 + 解码器」,输入仍是x_dict/edge_index_dict,线上服务把特征拼好直接喂入即可;如果线上只迭代编码器,也可以单独导出编码器复用既有服务,省一次全量发布。

五、上线前检查清单

  • RandomLinkSplitrev_edge_types是否与正向边一一对应,反向边没有留在训练集里
  • 时序 loader 的edge_label_time是否减过 1,测试边不存在于任何训练批次的时间窗之后
  • 对外只报 test split 的 RMSE/MAE,验证集数字仅作早停依据
  • Partitioner落盘后核对每个分片的graph.ptnode_feats.pt齐全,META.json与机器数一致
  • torch.jit.load在干净环境跑通一次端到端推理,输入输出 shape 与训练日志一致
  • neg_sampling比例已在离线网格上调过,而非沿用默认值

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

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

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

AI Agent 搜索 MCP 升级路径:从工具组合到基建体系

AI Agent接入外部工具与数据的标准化通信协议是Model &#xff08;MCP&#xff09;, 它已然成为智能体扩展能力的通用技术路径。早期的时候, Agent的搜索能力大多依靠单点的MCP工具组合得以实现, 以此适配原型验证以及轻量化场景。后来, 随着智能体开始逐渐向生产级场景落地, 体…

作者头像 李华
网站建设 2026/9/8 16:16:28

SSD存储接口深度剖析:从物理选型到固件开发全链路

1. 从一次加装硬盘的困惑说起&#xff1a;接口根本不是“接口”&#xff0c;是整套链路前几天有个朋友问我&#xff0c;说新加了一块固态硬盘&#xff0c;能不能把原来D盘的东西直接搬到SSD上去&#xff1f;我反问他SSD买了什么型号、装在哪个接口上&#xff0c;他一脸茫然回了…

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

FPGA上板调通测试实战指南:从仿真到板级稳定运行的关键方法

搞数字逻辑实验的同学&#xff0c;大概率都有过这种体验&#xff1a;仿真波形怎么测怎么对&#xff0c;逻辑功能挑不出毛病&#xff0c;结果代码一下到板子上&#xff0c;LED死活不亮&#xff0c;数码管乱跳&#xff0c;或者输出波形跟预期完全对不上。这时候最容易怀疑人生&am…

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

十周刨根问底Triton:不啃龙书的编译器实战学习路线

写这篇东西的起因其实挺简单&#xff1a;团队里来了个新人&#xff0c;深度学习跑得很溜&#xff0c;但一提到编译器就发怵&#xff0c;桌上那本红黑封面的“龙书”翻了两个月还在第2章&#xff0c;整天被词法分析和正则表达式按在地上摩擦。我说你别死磕了&#xff0c;换个思路…

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

AI Agent 搜索 MCP 选型指南:能力栈模型与基建核心

在 Model &#xff08;MCP&#xff09;变为 AI Agent 外部工具接入的通用标准之际, 搜索能力身为智能体的核心外部感知入口, 其被部署的形态正从单点工具插件, 朝着体系化的能力基建进行演进。当下, 多数以 MCP 推荐内容为主的情况, 乃是平级罗列工具, 缺少具备体系化的能力部署…

作者头像 李华