ARTICLE DETAIL

资讯详情

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

多任务GNN在供应链需求预测与风险评估中的落地实践

多任务GNN在供应链需求预测与风险评估中的落地实践 简介这份资源面向供应链管理研究者、数据科学家及行业从业者聚焦图神经网络在供应链优化中的落地应用。内容围绕多任务GNN模型展开涵盖供应链图结构建模、需求预测、风险评估与异常检测等核心任务并给出基于PyTorch Geometric的完整Python实现包括异构图数据构建、GCN/GAT/GraphSAGE模型定义、训练与评估流程。资源包内含1个PDF文件大小约708KB以论文与代码解释为主便于读者理解理论推导并复现实验。文中还提供来自孟加拉国快速消费品公司的多视角真实基准数据集在6类供应链分析任务上对比多种先进GNN模型结果显示其性能较传统方法提升10%至40%。目前已有105人学习适合希望掌握图结构建模、多任务学习及供应链韧性分析的读者参考。1. 供应链里的需求预测和风险评估为什么值得用多任务 GNN 重做一遍一个典型的供应链场景某区域仓管着三千个 SKU上游有几十家供应商下游覆盖几百个门店。需求预测团队用 LightGBM 跑时间序列风险评估团队用逻辑回归算供应商违约概率两套模型各自维护特征、各自调参、各自上线。问题出在两边用的特征高度重叠——供应商的历史交付延迟、区域仓的库存周转、门店的促销日历这些信号既影响需求也影响风险却被拆成两份互不相通的数据管道。更麻烦的是供应链天然是一张图SKU 之间有替代关系门店之间有调拨关系供应商之间有产能竞争关系。传统模型把每个节点当独立样本等于主动扔掉了这张图里最有价值的结构信息。多任务 GNN 要解决的就是这件事用一张异构图同时表达 SKU、门店、供应商、区域仓之间的多种关系让需求预测和风险评估共享底层图表示再通过任务头分别输出。共享表示带来的收益不是理论上的——当某个供应商出现交付波动这个信号会通过图结构传播到受影响的 SKU 需求预测上反之亦然。适合谁读已经跑通过单任务 GNN、想往多任务和供应链落地走的算法工程师以及做供应链数字化、需要判断这套方案边界的技术负责人。2. 多任务 GNN 的图结构设计与特征工程2.1 供应链异构图的三类节点和四类边把供应链抽象成异构图节点类型和边类型的设计直接决定模型能学到什么。常见做法是定义三类核心节点SKU 节点、门店节点、供应商节点区域仓可以作为 SKU 节点的属性而不是独立节点避免图过于稀疏。边类型至少覆盖四类边类型起点终点边特征示例供应关系供应商SKU供货比例、账期、历史准时率销售关系SKU门店近 30 天销量、铺货率替代关系SKUSKU品类相似度、价格带重叠度调拨关系门店门店调拨频次、平均调拨时长替代关系和调拨关系是很多团队容易漏掉的。没有替代边模型无法理解“A 缺货时 B 的销量会涨”这种信号没有调拨边门店之间的库存联动就丢了。边特征不要只放静态属性把近 7 天、近 30 天的滚动统计量拼进去模型才能感知动态变化。2.2 节点特征里必须有的时间窗口统计量节点特征分两类静态属性和动态统计量。静态属性比如 SKU 的品类、价格带、供应商的合作年限。动态统计量是重点建议至少构造三个时间窗口import pandas as pd import numpy as np def build_node_features(sales_df, window_days[7, 14, 30]): 为每个 SKU-门店组合构造多时间窗口统计特征 sales_df: 包含 date, sku_id, store_id, sales_qty, stock_qty 的长表 features [] for w in window_days: # 滚动销量均值和标准差 roll sales_df.groupby([sku_id, store_id])[sales_qty].transform( lambda x: x.rolling(w, min_periods1).mean() ) roll_std sales_df.groupby([sku_id, store_id])[sales_qty].transform( lambda x: x.rolling(w, min_periods1).std() ) # 库存周转天数当前库存 / 窗口内日均销量 avg_daily sales_df.groupby([sku_id, store_id])[sales_qty].transform( lambda x: x.rolling(w, min_periods1).mean() ) turnover sales_df[stock_qty] / (avg_daily 1e-6) features.append(pd.DataFrame({ fsales_mean_{w}d: roll, fsales_std_{w}d: roll_std, fturnover_{w}d: turnover })) return pd.concat(features, axis1)这段代码的关键点min_periods1保证序列开头不会产生 NaNturnover分母加1e-6防止除零三个窗口分别捕捉短期波动、中期趋势和月度节奏。实际落地时还要加一个“促销标记”的布尔特征否则促销期的销量尖峰会被模型当成噪声。2.3 用 PyG 搭一个可跑通的异构图图神经网络代码的入口通常是把上面的 DataFrame 转成 PyTorch Geometric 的HeteroData对象。下面是最小可运行版本import torch from torch_geometric.data import HeteroData def build_hetero_graph(sku_feat, store_feat, supplier_feat, supply_edge, sell_edge, subst_edge, transfer_edge): 构建供应链异构图 sku_feat: (num_sku, d_sku) 张量 supply_edge: (2, num_supply) 供应商-SKU 的边索引 data HeteroData() data[sku].x torch.tensor(sku_feat, dtypetorch.float) data[store].x torch.tensor(store_feat, dtypetorch.float) data[supplier].x torch.tensor(supplier_feat, dtypetorch.float) data[supplier, supplies, sku].edge_index torch.tensor(supply_edge, dtypetorch.long) data[sku, sells, store].edge_index torch.tensor(sell_edge, dtypetorch.long) data[sku, substitutes, sku].edge_index torch.tensor(subst_edge, dtypetorch.long) data[store, transfers, store].edge_index torch.tensor(transfer_edge, dtypetorch.long) return dataHeteroData的好处是每种边类型独立维护消息传递路径supplies边和substitutes边不会互相干扰。注意substitutes是 SKU 到 SKU 的自环类型边PyG 支持这种同类型节点间的异质边。如果 SKU 数量超过十万edge_index用torch.long存会占不少内存可以考虑按品类分片建图。3. 多任务 GNN 模型实现共享编码器加双任务头3.1 共享图编码器用 GraphSAGE 还是 GAT选型上供应链图的特点是边类型多、节点度数差异大热门 SKU 可能连几百个门店冷门 SKU 只连几个。GraphSAGE 的邻居采样机制对高度数节点更友好训练时不会因为某个 SKU 的邻居太多而爆显存。GAT 的注意力权重可解释性更好但计算开销随度数平方增长。我一般会先用 GraphSAGE 跑通 baseline确认多任务共享表示有效之后再在关键边类型上换 GAT 做对比。下面是一个两层 HeteroGraphSAGE 编码器import torch.nn as nn from torch_geometric.nn import SAGEConv, HeteroConv class SupplyChainEncoder(nn.Module): def __init__(self, hidden_dim64, out_dim32): super().__init__() # 第一层每种边类型独立聚合 self.conv1 HeteroConv({ (supplier, supplies, sku): SAGEConv((-1, -1), hidden_dim), (sku, sells, store): SAGEConv((-1, -1), hidden_dim), (sku, substitutes, sku): SAGEConv((-1, -1), hidden_dim), (store, transfers, store): SAGEConv((-1, -1), hidden_dim), }, aggrsum) # 第二层输出统一维度的节点表示 self.conv2 HeteroConv({ (supplier, supplies, sku): SAGEConv((-1, -1), out_dim), (sku, sells, store): SAGEConv((-1, -1), out_dim), (sku, substitutes, sku): SAGEConv((-1, -1), out_dim), (store, transfers, store): SAGEConv((-1, -1), out_dim), }, aggrsum) self.norm nn.LayerNorm(out_dim) def forward(self, x_dict, edge_index_dict): x_dict self.conv1(x_dict, edge_index_dict) x_dict {k: torch.relu(v) for k, v in x_dict.items()} x_dict self.conv2(x_dict, edge_index_dict) x_dict {k: self.norm(v) for k, v in x_dict.items()} return x_dictHeteroConv的aggrsum表示同一目标节点从不同边类型收到的消息直接相加。如果某类边信号特别强可以改成aggrmean或自定义加权。SAGEConv((-1, -1), hidden_dim)里的-1让 PyG 自动推断输入维度省去手动计算每种节点特征维度。3.2 需求预测头分位数回归比 MSE 更实用需求预测任务如果只用 MSE 损失模型会倾向于预测均值但供应链补货决策关心的是“95 分位需求是多少”。用分位数回归头一次输出 P50、P90、P95 三个分位点class DemandHead(nn.Module): def __init__(self, in_dim, quantiles[0.5, 0.9, 0.95]): super().__init__() self.quantiles quantiles self.fc nn.Sequential( nn.Linear(in_dim, 64), nn.ReLU(), nn.Linear(64, len(quantiles)) ) def forward(self, sku_emb): return self.fc(sku_emb) # (num_sku, num_quantiles) def quantile_loss(pred, target, quantiles): pinball loss对每个分位点分别计算 losses [] for i, q in enumerate(quantiles): err target - pred[:, i] losses.append(torch.max((q - 1) * err, q * err)) return torch.stack(losses, dim1).mean()pinball loss的逻辑当真实值大于预测值时高估的惩罚权重是q低估的惩罚权重是1-q。对 P95 来说低估的惩罚远大于高估模型自然会把预测值往上推。这比 MSE 更贴合补货场景的安全库存逻辑。3.3 风险评估头二分类加 focal loss 处理样本不均衡供应商风险事件延迟交付、质量事故在真实数据里占比通常不到 5%标准交叉熵会被多数类主导。用 focal loss 让模型聚焦难分样本class RiskHead(nn.Module): def __init__(self, in_dim): super().__init__() self.fc nn.Sequential( nn.Linear(in_dim, 32), nn.ReLU(), nn.Linear(32, 1) ) def forward(self, supplier_emb): return self.fc(supplier_emb).squeeze(-1) def focal_loss(logits, targets, alpha0.75, gamma2.0): alpha 控制正负样本权重gamma 控制难易样本聚焦程度 bce nn.functional.binary_cross_entropy_with_logits( logits, targets, reductionnone ) pt torch.exp(-bce) loss alpha * (1 - pt) ** gamma * bce return loss.mean()alpha0.75表示正样本有风险权重更高gamma2.0是常见默认值。如果风险事件占比极低低于 1%可以把alpha提到 0.9同时监控验证集上的召回率而不是准确率。3.4 多任务损失加权不确定性加权比网格搜索省事两个任务的损失量级不同直接相加会让需求预测主导梯度。常见做法是手动调权重但更省事的是用 Kendall 的不确定性加权class MultiTaskLoss(nn.Module): def __init__(self): super().__init__() # 可学习的对数方差参数 self.log_sigma_demand nn.Parameter(torch.zeros(1)) self.log_sigma_risk nn.Parameter(torch.zeros(1)) def forward(self, demand_loss, risk_loss): # 精度 1 / sigma^2用 log_sigma 保证数值稳定 precision_d torch.exp(-self.log_sigma_demand) precision_r torch.exp(-self.log_sigma_risk) total (precision_d * demand_loss self.log_sigma_demand precision_r * risk_loss self.log_sigma_risk) return total.mean()训练过程中如果某个任务的损失下降困难对应的log_sigma会增大该任务的权重自动降低。这比手动设0.7 * demand_loss 0.3 * risk_loss更适应训练动态。注意这两个参数要用更大的学习率比如主网络的 10 倍否则收敛太慢。4. 训练、验证与线上推理的工程细节4.1 按时间切分数据集别用随机切分供应链数据有时间泄漏风险如果用随机切分模型可能看到未来数据来预测过去。正确做法是按时间切分——用前 8 个月训练第 9 个月验证第 10 个月测试。图结构也要对应切分验证集和测试集的边不能出现在训练图里否则消息传递会泄漏未来信息。def temporal_split(edge_df, train_end, val_end): 按时间戳切分边返回三个子图 train_edges edge_df[edge_df[date] train_end] val_edges edge_df[(edge_df[date] train_end) (edge_df[date] val_end)] test_edges edge_df[edge_df[date] val_end] return train_edges, val_edges, test_edges实际落地时验证集和测试集的节点特征可以用全量数据构造因为特征本身不涉及标签但边必须严格按时间切。如果某个 SKU 在验证期才首次出现它的图邻居可能为空这时候要回退到只用节点特征的 baseline 预测。4.2 训练循环里的梯度裁剪和早停多任务 GNN 的梯度容易在共享编码器里累积放大尤其是图卷积层数超过 3 层时。训练循环里加梯度裁剪optimizer torch.optim.Adam(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience5) for epoch in range(100): model.train() optimizer.zero_grad() x_dict encoder(data.x_dict, data.edge_index_dict) demand_pred demand_head(x_dict[sku]) risk_pred risk_head(x_dict[supplier]) d_loss quantile_loss(demand_pred[train_sku_mask], demand_target[train_sku_mask], [0.5, 0.9, 0.95]) r_loss focal_loss(risk_pred[train_supplier_mask], risk_target[train_supplier_mask]) loss mtl_loss(d_loss, r_loss) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() # 验证集评估 val_metric evaluate(model, val_data) scheduler.step(val_metric) if early_stop(val_metric, patience10): breakclip_grad_norm_的max_norm1.0是经验值如果训练不稳定可以降到 0.5。ReduceLROnPlateau在验证指标不再下降时自动降学习率配合早停能省不少调参时间。4.3 线上推理子图采样比全图推理更现实全图推理在 SKU 数量大时延迟不可接受。线上服务通常只对当天有变化的子图做推理——比如某个供应商出了交付异常只重新计算受影响的 SKU 和门店。实现上用torch_geometric.loader.NeighborLoader做 k 跳采样from torch_geometric.loader import NeighborLoader # 只对需要预测的 SKU 节点做 2 跳采样 loader NeighborLoader( data, num_neighbors{key: [10, 5] for key in data.edge_types}, batch_size256, input_nodes(sku, predict_mask), shuffleFalse ) model.eval() with torch.no_grad(): for batch in loader: x_dict encoder(batch.x_dict, batch.edge_index_dict) preds demand_head(x_dict[sku][:batch[sku].batch_size])num_neighbors里[10, 5]表示第一跳采样 10 个邻居第二跳采样 5 个。input_nodes指定只对predict_mask为 True 的 SKU 做推理。这样单次推理只涉及几百个节点延迟可以控制在 50ms 以内。5. 多任务 GNN 的调参技巧与效果验证5.1 共享编码器层数2 层是供应链图的甜点区图卷积层数不是越多越好。供应链图的直径通常不大供应商到门店一般 3 跳以内2 层编码器已经能覆盖大部分消息传递路径。层数加到 4 层以上会出现过平滑——所有 SKU 的表示趋同需求预测的区分度反而下降。验证方法很简单训练完后计算 SKU 表示的 pairwise 余弦相似度均值如果超过 0.9说明过平滑了该减层。5.2 任务权重的手动兜底方案不确定性加权虽然省事但训练初期两个log_sigma参数随机初始化可能导致某个任务在前几个 epoch 完全没学到东西。兜底方案是前 10 个 epoch 用固定权重比如需求 0.6、风险 0.4之后再切换到可学习权重。切换时机可以通过观察两个任务的验证损失是否都开始下降来判断。5.3 效果验证需求预测看 WAPE风险评估看 PR-AUC需求预测的评估指标建议用 WAPE加权绝对百分比误差比 MAPE 对低销量 SKU 更友好def wape(y_true, y_pred): 加权绝对百分比误差分母用总销量而非单点销量 return np.sum(np.abs(y_true - y_pred)) / np.sum(np.abs(y_true))风险评估用 PR-AUC 而不是 ROC-AUC因为风险事件是少数类ROC-AUC 会被大量真负例拉高。多任务模型相比单任务 baseline需求预测的 WAPE 通常能降 3 到 8 个百分点风险评估的 PR-AUC 能提升 5 到 12 个百分点。如果提升不明显先检查图结构里替代边和调拨边是否真的建对了——这两类边往往是多任务共享表示能否生效的关键。本文还有配套的精品资源点击获取
返回列表