
PyG 供应链运输成本预测指南HeteroData 建图到 jit 部署的完整流程【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric本文基于 PyTorch GeometricPyG以预测「仓库→客户」运输成本为任务完整走一遍 HeteroData 异构图建图、SAGEConv 边级回归、时序采样防泄漏、分布式两级扩展与 torch.jit 部署代码可直接改造为你的物流网络。三句话说明本文解决的是「线路成本沿关系传导、单表模型看不到」的预测难题你会得到一条从建图、切分、训练、评估到上线的可运行管线适合有 Python/PyTorch 基础、要给自己的物流网络做预测服务的工程师。1. 为什么一条运输线路的成本不好预测报价单上的成本只和距离、车型有关但真实成本受上下游联动供应商产能波动影响仓库库存库存变化后改走其他运输线最终影响客户交付。这类传导发生在表与表之间用单表特征模型只能靠人工拼交叉特征去近似。图方法的做法是把实体和关系放进同一个结构让邻域消息机制把上下游状态聚合进节点表示再用边两端点的表示预测这条线路的成本。后文所有代码基于 PyG 2.x任务统一为「预测 warehouse→customer 边的成本」。2. HeteroData 建图步骤供应链里有供应商、仓库、客户、产品四类实体和多种关系HeteroData用「节点类型 边类型」双索引直接表达不需要合并成一张大节点表。import torch from torch_geometric.data import HeteroData g HeteroData() # 节点特征取自 ERP/WMS 现成字段先做 z-score 统一量纲 g[supplier].x torch.randn(120, 8) # 产能、区位、历史履约率 g[warehouse].x torch.randn(30, 8) # 库容、周转天数、租金 g[customer].x torch.randn(5000, 8) # 下单频次、账期、区域 g[product].x torch.randn(300, 8) # 体积重、温层、单价 # edge_index 形状为 2xK第 0 行是起点节点第 1 行是终点节点 g[supplier, supplies, warehouse].edge_index sup_wh g[warehouse, stores, product].edge_index wh_prod g[warehouse, ships_to, customer].edge_index wh_cust三个容易踩的点边类型只保留有业务语义的组合不穷举没有特征的纯关系型节点先用独热 ID 占位数据补齐后再替换z-score 按字段离线做避免「库容吨」和「租金元/平米」这类单位差异在消息传递中互相压制。仓库里的 examples/hetero/hetero_link_pred.py 是完整的异构图边级回归示例该例是用户-电影评分换成供应链实体后结构不变。3. 先切分RandomLinkSplit 与反向边泄漏边级任务切的是边集而不是节点集。一个典型的坑忘记处理反向边于是测试边的反向边留在了训练集里训练时模型已经「见过」测试线路test 指标虚高、上线即失效。from torch_geometric.transforms import RandomLinkSplit train_g, val_g, test_g RandomLinkSplit( num_val0.1, num_test0.1, neg_sampling_ratio0.0, # 回归任务不需要负样本 edge_types[(warehouse, ships_to, customer)], rev_edge_types[(customer, rev_ships_to, warehouse)], )(g)rev_edge_types的作用是把被切走的边的反向边同步挪进对应切分保证训练集与测试集之间没有交叉。回归任务把neg_sampling_ratio置零避免负样本干扰目标值。4. 模型与评估SAGEConv to_heteroRMSE/MAE 对账金额模型结构分两段编码器把每个节点编码成向量解码器把线路两端点的向量拼接后过小 MLP 输出标量。编码器按同质 GNN 写交给to_hetero按元数据自动展开成每种边类型一组独立参数。from torch_geometric.nn import SAGEConv, to_hetero class Encoder(torch.nn.Module): def __init__(self, dim): super().__init__() # -1 表示输入维度自动推断让 to_hetero 完成异质展开 self.conv1 SAGEConv((-1, -1), dim) self.conv2 SAGEConv((-1, -1), dim) def forward(self, x, edge_index): return self.conv2(self.conv1(x, edge_index).relu()) class PairDecoder(torch.nn.Module): def __init__(self, dim): super().__init__() self.net torch.nn.Sequential( torch.nn.Linear(2 * dim, dim), torch.nn.ReLU(), torch.nn.Linear(dim, 1), ) def forward(self, z_dict, edge_label_index): src, dst edge_label_index pair torch.cat([z_dict[warehouse][src], z_dict[customer][dst]], dim-1) return self.net(pair).squeeze(-1) class CostModel(torch.nn.Module): def __init__(self, dim): super().__init__() self.encoder to_hetero(Encoder(dim), g.metadata(), aggrsum) self.decoder PairDecoder(dim) def forward(self, x_dict, edge_index_dict, edge_label_index): return self.decoder(self.encoder(x_dict, edge_index_dict), edge_label_index)两句点评不写-1就得为每种节点类型手工指定输入维度自动推断是to_hetero发挥价值的前提解码器只做拼接与小 MLP不必一上来引入注意力。训练用 MSE 即可汇报指标用 RMSE 和 MAE 双口径因为它们能直接换算成金额。import torch.nn.functional as F model CostModel(64) opt torch.optim.Adam(model.parameters(), lr0.01) ROUTE (warehouse, ships_to, customer) def train_one_epoch(): model.train() opt.zero_grad() pred model(train_g.x_dict, train_g.edge_index_dict, train_g[ROUTE].edge_label_index) loss F.mse_loss(pred, train_g[ROUTE].edge_label) loss.backward() opt.step() return float(loss) torch.no_grad() def evaluate(split): model.eval() pred model(split.x_dict, split.edge_index_dict, split[ROUTE].edge_label_index) label split[ROUTE].edge_label.float() return float(F.mse_loss(pred, label).sqrt()), \ float(F.l1_loss(pred, label)) for epoch in range(1, 201): train_one_epoch() print(epoch, evaluate(val_g))数字怎么对上业务设成本单位为元/箱test 集 MAE850 元、月均 6000 箱则模型与真实成本的总偏差约为 850 x 6000 ≈ 510 万元/月这个量级就是和财务对齐「预期对账差额」的依据如果现行人工报价系统性高于实际成本模型带来的改善就是两者之差。记住口径val 只用于早停模型好坏以 test split 为准。5. 线路天天变LinkNeighborLoader 时序采样防泄漏 ⏱️静态切边对「用上个季度的数据」够用但运输线路每天都在增删上周刚开的新线路不该出现在训练子图里。解法在 examples/hetero/recommender_system.py 中给每条边带时间戳采样时用目标边的预测时点约束每一跳的邻居截断。from torch_geometric.loader import LinkNeighborLoader loader LinkNeighborLoader( datag, num_neighbors[5, 5], edge_label_index((warehouse, ships_to, customer), train_edge_index), edge_label_timetrain_time - 1, # 关键减 1 time_attrtime, temporal_strategylast, batch_size256, shuffleTrue, )两处参数各管一件事edge_label_time减 1 是因为采样按「边时间戳早于该时刻」截断不减就可能把预测时点的那条边本身采进子图形成未来泄漏temporal_strategylast让每一跳只取截断时刻之前的最后若干邻居。若任务升级为「预测下一条线路」给 loader 加neg_samplingdict(modebinary, amount2)构造负样本并用torch_geometric.metrics的LinkPredPrecision/LinkPredRecall评估Precision10 的读法是「给每条线路推 10 个候选合作方平均命中 K 个真实发生过往来的」。6. 图大到单机装不下Partitioner DistNeighborLoader 两级扩展为什么必须扩展大客户订单图边数可达上亿整图邻接和特征都放不下内存。torch_geometric/distributed/把扩展拆成两级。离线切图Partitioner把节点、边和特征按分片落盘每个partN/目录下是graph.pt与node_feats.pt。底层 METIS 切分是非确定性的所以分区文件在一台机器上生成后分发到整个集群保证各进程读同一份分区。在线采样DistNeighborLoader绑定本进程分片本地邻居直接读跨分片邻居走 RPC 从远端进程拉取。采样开销从「全图」降到「本机分片 一跳远程」训练吞吐随机器数近似线性扩展。from torch_geometric.distributed import ( DistNeighborLoader, LocalFeatureStore, LocalGraphStore, Partitioner, ) # 第一步离线执行一次产出 partitions/ 目录 Partitioner(datag, num_parts8, root./partitions) # 第二步每个训练进程只加载自己的分片然后开始采样 feature_store, graph_store LocalFeatureStore(ctx), LocalGraphStore(ctx) feature_store.load_data(./partitions, supply_chain) graph_store.load_data(./partitions, supply_chain) loader DistNeighborLoader( data(feature_store, graph_store), num_neighbors[10, 10], master_addr10.0.0.1, master_port8901, current_ctxctx, )采样器接口与单机NeighborLoader一致训练循环无需改动只替换数据来源。7. torch.jit 导出 GNN 完成部署训练完用torch.jit.script脚本化导出线上推理侧不再依赖 Python 训练环境做法同 examples/jit/gin.pyscripted torch.jit.script(model) torch.jit.save(scripted, route_cost.pt) served torch.jit.load(route_cost.pt) pred served(x_dict, edge_index_dict, edge_label_index)导出的是「编码器 解码器」整体输入仍为x_dict/edge_index_dict线上把特征拼装好直接喂入即可若线上只更新编码器特征分布变了但解码逻辑不变也可以只导出编码器单独服务缩短发布链路。8. 避坑清单与下一步 以下每个坑都会让一整轮训练失效RandomLinkSplit漏写rev_edge_types测试边的反向边留在训练集test 指标虚高。时序采样忘了edge_label_time - 1子图里采到未来边。用 val 数字判断模型好坏只有 test split 是结论依据。跨量纲特征未归一化消息传递被大尺度字段主导。回归任务保留默认负采样比例负样本干扰边级目标值。接下来按价值排序可以做的事多任务头同一份z_dict上同时预测成本、时效、断供概率不同 decoder 共享编码器一次采样服务三个任务。链路预测任务重点调负采样比例它对 PrecisionK 的影响通常大于模型层数。把 MSE 换成分段线性损失压低高价值大客户线路上的偏差业务侧更关心这段而不是整体 MAE。【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考