
把供应链变成一张异构图用 PyTorch Geometric 预测运输成本到底能省多少钱【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric仓库到客户这一单运过去要花多少这个问题的答案通常散落在 ERP 的几张表里单张表算不出。本文用 PyTorch GeometricPyG图神经网络库把供应链组织成异构图预测「仓库 → 客户」边上的运输成本从建图、建模、时序采样一直走到分布式扩展与上线部署。速览用HeteroData承载 4 类节点、3 类边的供应链异构图边级回归直接预测线路成本SAGEConv写同质编码器交给to_hetero自动展开成异质模型MSE 训练LinkNeighborLoader按时间戳采样防未来泄漏torch_geometric/distributed/扛大图torch.jit扛部署适合读者有 Python/PyTorch 基础想给自己的物流网络做预测或推荐的工程师。读完你可以复现一套「建图 → 边级预测 → 时序训练 → 部署」的完整链路。1 为什么表格算不动影响传导供应链运输成本预测的问题定义传统做法是供应商、仓库、客户各建各的表各查各的。但供应链里影响是沿着关系传递的某供应商产能掉点 → 对应仓库缺货 → 线路改走更贵的备选线 → 客户收货延期。这一棒接力跨过多张表单表里你只看到结果看不到传导过程几张表 join 起来口径又难对齐。图的自然表达方式是把实体与关系写进同一个数据结构让特征影响像接力棒一样沿边逐跳传递关系传了几跳、模型就能解释几跳。本文的目标就是在这种异构图上锁定一个业务数字——每条「仓库 → 客户」线路的运输成本。2 HeteroData 构建供应商/仓库/客户/产品装进同一容器HeteroData是 PyG 的异构图数据容器节点类型存在data[node_type].x边类型存在data[src, rel, dst].edge_index。节点特征直接从 ERP/WMS 取现成字段统一做z-score归一化个别纯关系型节点没有特征可以先用独热 ID 顶上去。边类型不必穷举所有组合只保留有业务语义的几条节点类型示例规模特征来源supplier 供应商120产能、区位、历史履约率warehouse 仓库30库容、周转天数、租金customer 客户5000下单频次、账期、区域product 产品300体积重、温层、单价建图代码很短边索引是关键——每条边用 2×E 的矩阵表达第 0 行是起点节点编号第 1 行是终点import torch from torch_geometric.data import HeteroData data HeteroData() data[supplier].x torch.randn(120, 8) data[warehouse].x torch.randn(30, 8) data[customer].x torch.randn(5000, 8) data[product].x torch.randn(300, 8) # edge_index 为 2×E第 0 行起点第 1 行终点 data[supplier, supplies, warehouse].edge_index sup_wh_idx data[warehouse, stores, product].edge_index wh_prod_idx data[warehouse, transports, customer].edge_index wh_cust_idx仓库自带示例examples/hetero/hetero_link_pred.py用「用户-评分-电影」异构图做同类任务把节点名换成供应链实体后结构完全通用。3 SAGEConv 边级预测to_hetero 异质展开自动建模型边级回归的思路很直白先把每个节点编码成向量再取边两端点的向量拼起来过一个小 MLP 输出标量成本。编码器按「同质 GNN」来写——SAGEConv((-1, -1), ...)中的-1表示输入维度交由数据推断to_hetero依赖这个约定才能给每种边类型自动配独立参数不用你手写一堆 conv。建模前先切边。RandomLinkSplit负责把目标边拆成训练/验证/测试三份注意rev_edge_types必须带上反向边否则反向边会漏进训练集from torch_geometric.nn import SAGEConv, to_hetero from torch_geometric.transforms import RandomLinkSplit train_data, val_data, test_data RandomLinkSplit( num_val0.1, num_test0.1, neg_sampling_ratio0.0, # 回归不需要负样本 edge_types[(warehouse, transports, customer)], rev_edge_types[(customer, rev_transports, warehouse)], )(data) class GNNEncoder(torch.nn.Module): def __init__(self, hidden_channels, out_channels): super().__init__() self.conv1 SAGEConv((-1, -1), hidden_channels) self.conv2 SAGEConv((-1, -1), out_channels) def forward(self, x, edge_index): return self.conv2(self.conv1(x, edge_index).relu())解码器按edge_label_index取出两端点向量拼接后输出标量最外层的Model用to_hetero一行完成异质展开class EdgeDecoder(torch.nn.Module): def __init__(self, hidden_channels): super().__init__() self.lin1 torch.nn.Linear(2 * hidden_channels, hidden_channels) self.lin2 torch.nn.Linear(hidden_channels, 1) def forward(self, z_dict, edge_label_index): row, col edge_label_index z torch.cat([z_dict[warehouse][row], z_dict[customer][col]], dim-1) return self.lin2(self.lin1(z).relu()).view(-1) class Model(torch.nn.Module): def __init__(self, hidden_channels): super().__init__() self.encoder to_hetero( GNNEncoder(hidden_channels, hidden_channels), metadatadata.metadata(), aggrsum) self.decoder EdgeDecoder(hidden_channels) def forward(self, x_dict, edge_index_dict, edge_label_index): z_dict self.encoder(x_dict, edge_index_dict) return self.decoder(z_dict, edge_label_index)metadatadata.metadata()提供「节点类型 边类型」清单aggrsum决定消息聚合方式两者均与官方示例保持一致。4 MSE 训练循环把 RMSE/MAE 换算成业务金额训练就是最普通的回归循环损失用 MSEimport torch.nn.functional as F model Model(64) optimizer torch.optim.Adam(model.parameters(), lr0.01) def train(): model.train() optimizer.zero_grad() pred model(train_data.x_dict, train_data.edge_index_dict, train_data[warehouse, transports, customer].edge_label_index) loss F.mse_loss(pred, train_data[warehouse, transports, customer].edge_label) loss.backward() optimizer.step() return float(loss) torch.no_grad() def test(data): model.eval() et (warehouse, transports, customer) pred model(data.x_dict, data.edge_index_dict, data[et].edge_label_index) target data[et].edge_label.float() return float(F.mse_loss(pred, target).sqrt()), float(F.l1_loss(pred, target)) for epoch in range(1, 201): print(train(), test(val_data))评估取两个数RMSEF.mse_loss(...).sqrt()与 MAEF.l1_loss两者都能直接换成钱。举个例子设edge_label单位是千元/单测试集 MAE ≈ 0.12 千元/单即单均偏差 120 元月单量 5 万单则月度平均偏差 ≈ 0.12 × 50000 600 万元。这个数字可以直接和现行手工维护的固定报价对比——如果报价偏差比它更大模型就有上线的价值。再强调一句判断好坏只看 test splitval 上的数字只配用来早停。5 LinkNeighborLoader 时序采样堵住未来信息泄漏上一节的静态切边有个隐患运输关系每天在变上周新开的线路不该出现在历史训练里。仓库示例examples/hetero/recommender_system.py的做法是按时间戳切分给边打上时间戳采样时只回看「预测时点」之前的历史——类似地铁换乘只数已经到站的列车未来的班次不算。from torch_geometric.loader import LinkNeighborLoader loader LinkNeighborLoader( datadata, num_neighbors[5, 5], edge_label_index((warehouse, transports, customer), edge_index), edge_label_timeedge_time - 1, # -1不采到目标边本身 time_attrtime, temporal_strategylast, # 每跳只取截断时刻之前的邻居 batch_size256, shuffleTrue, )关键参数作用edge_label_timeedge_time - 1给每条待预测边设定采样截止时刻偏移 -1 避免采到目标边本身temporal_strategylast每一跳邻居采样都受截止时刻约束只取时刻之前的边neg_samplingdict(modebinary, amount2)链路预测场景下每条正例配 2 条负样本两个参数配合后每一跳采样都被锁在「预测时点」之前机制上杜绝未来边进入子图——这类泄漏在物流场景里比模型本身更常出错。如果任务从回归升级为链路预测判断线路是否存在评估换成torch_geometric.metrics的LinkPredPrecision(k)/LinkPredRecall(k)Precision10 回答「每条线路推荐的 10 个候选合作方里平均有几个真实发生过往来」Recall10 回答「真实往来被推荐列表覆盖了多少」。6 GNN 分布式训练与 torch.jit 模型部署节点过百万、单机装不下时torch_geometric/distributed/提供两级扩展。离线阶段Partitioner把节点与特征按分片落盘每个partN/目录下是graph.pt与node_feats.pt在线阶段DistNeighborLoader绑定本机分片本地邻居直接读盘跨分片邻居走 RPC 从远端机器拉取。采样开销从「全图」降为「本机分片 一跳远程」训练吞吐随机器数近似线性扩展对大客户订单边动辄上亿的供应链网络这一步基本是必选项。训练完的模型用torch.jit.script脚本化导出做法参考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线上服务把特征拼好直接喂入即可如果线上只更新编码器特征变了、解码关系不变也可以只导编码器单独服务。7 要点清单与高频坑位要点清单建图HeteroData承载多节点/多边类型边用 2×E 索引节点特征做z-score归一化模型SAGEConv((-1, -1), ...)维度自推断to_hetero按 metadata 展开成异质模型边类型各自独立参数边级回归两端点向量拼接 → MLP 输出标量MSE 训练RMSE/MAE 对账到业务金额时序采样LinkNeighborLoadertemporal_strategylastedge_label_time - 1防未来泄漏扩展与部署torch_geometric/distributed/负责切图与跨机采样torch.jit负责脚本化上线下一步按优先级多任务解码头同一个z_dict上同时预测成本、时效、断供概率不同 decoder 共享编码器链路预测调neg_sampling比例负样本量对 PrecisionK 影响很大把 MAE 换成业务可解释的损失如分段线性让高价值「大客户线路」偏差更小高频坑位反向边泄漏RandomLinkSplit漏写rev_edge_types反向边悄悄进训练集未来信息泄漏时序采样缺temporal_strategylast或edge_label_time偏移未来边混进子图拿 val 判断效果val 只用于早停拿它当效果指标会系统性高估【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考