十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

供应链里那些算不清的账,交给图神经网络:PyG 异构图运输成本预测实战

供应链里那些算不清的账,交给图神经网络:PyG 异构图运输成本预测实战 供应链里那些算不清的账交给图神经网络PyG 异构图运输成本预测实战【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric用 PyTorch GeometricPyG图神经网络库把供应商、仓库、客户、产品装进同一张异构图再预测「仓库→客户」这条边的运输成本从建模、时序采样到分布式扩展和部署本文走一遍完整链路。速览本文讲如何用 PyG 的HeteroData异构图数据容器SAGEConv编码器做边级成本回归并用LinkNeighborLoader处理随时间变化的运输关系。适合读者有 Python/PyTorch 基础、想给自己的物流网络做预测或推荐的工程师。1 把供应链装进一张图HeteroData 异构图数据构建传统做法是把供应链拆成一堆 Excel 和 SQL 视图分别分析但「供应商产能不足 → 某仓缺货 → 改走另一条运输线 → 客户延期」这类影响是沿着关系传递的表与表之间的传导在单表里根本看不到。图的价值就是把实体和它们的关系画在同一张纸上让消息沿着关系传播。用 PyG 的HeteroData建图关键是节点和边都用「类型 索引」表达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) # 体积重、温层、单价 # 边2x2 的 index第 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节点特征从 ERP/WMS 里取现成字段做z-score归一化即可特征缺的节点如纯关系型节点可以先用独热 ID 顶上去。边类型不用穷举所有组合只保留有业务语义的那几条。仓库自带示例 examples/hetero/hetero_link_pred.py 用的是「用户-评分-电影」异构图把节点名换成供应链实体后结构完全通用。2 边级回归SAGEConv 编码器 to_hetero 做运输成本预测边级预测预测某条线路的成本、时效、断供概率是供应链里最常见的落地点。思路先把每个节点编码成向量再取边两个端点的向量拼起来过一个小 MLP输出标量。下面三段是模型核心to_hetero负责把「同质 GNN」按元数据自动展开成异构图模型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) # -1: 输入维度自动推断 self.conv2 SAGEConv((-1, -1), out_channels) def forward(self, x, edge_index): return self.conv2(self.conv1(x, edge_index).relu()) 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)SAGEConv((-1, -1), ...)里的-1表示输入维度由数据推断这样to_hetero才能自动给每种边类型配独立参数。注意RandomLinkSplit时rev_edge_types要带上反向边否则反向边会泄漏到训练集。训练就是最普通的 MSE 回归评估函数同时算了 RMSE 和 MAE方便对到业务口径import 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))训练完看test(test_data)。RMSE 和 MAE 都能直接换算成钱测试集上 MAE0.8 千元/单乘以月单量 10 万就是「模型平均偏差 ≈ 80 万元/月」——这个数拿去和现在拍脑袋的固定报价比就知道该不该用模型出价。判断好坏只看 test splitval 上的数字只用来早停。3 动态运输网络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, )temporal_strategylast让每一跳采样都受edge_label_time约束只采到「预测时点」之前的历史边从机制上杜绝未来信息泄漏——这在物流场景里比模型本身更容易出错。评估侧换成推荐口径LinkPredPrecision(k)/LinkPredRecall(k)来自torch_geometric.metrics。Precision10 可以解读为「给每条线路推荐 10 个候选合作方平均有 K 个是真发生了往来的」召回率则回答「真实合作被推荐列表覆盖了多少」。链路预测任务记得在 loader 里加neg_samplingdict(modebinary, amount2)造负样本。4 图大到装不下内存分布式邻居采样当节点过百万单机装不下整张图torch_geometric/distributed/提供两级扩展离线切图Partitioner把节点和特征按分片落盘每个partN/下是graph.pt与node_feats.pt在线采样DistNeighborLoader绑定本分片本地邻居直接读跨分片邻居走 RPC 从远端拉。效果是采样开销从「全图」降到「本机分片 一跳远程」训练吞吐随机器数近似线性扩展。对供应链这类边数远超单机的网络大客户订单边动辄上亿这一步基本是必选项。5 上线部署torch.jit 脚本化导出与加载训练完的模型用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线上服务把特征拼装好直接喂入即可如果线上只更新编码器特征变了但解码关系不变也可以只导编码器单独服务。6 要点回顾与下一步要点回顾建模HeteroData表达多节点/多边类型边用 2×E 的 index节点特征做归一化模型SAGEConv((-1, -1), ...)让维度自动推断to_hetero按元数据展开成异构图边级回归拼接两端点向量 → 小 MLP 输出标量MSE 训练RMSE/MAE 对账到业务金额时序采样LinkNeighborLoadertemporal_strategylastedge_label_time - 1防未来泄漏扩展torch_geometric/distributed/做切图与跨机采样torch.jit脚本化部署下一步可以做的事按优先级给解码器加多任务头同一个z_dict上同时预测成本、时效、断供概率用不同的 decoder 共享编码器链路预测任务加neg_sampling调参负样本比例对 PrecisionK 影响很大把 MAE 换成业务可解释的损失如分段线性让模型在「大客户线路」上偏差更小【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表