库存预测这件事,做供应链的人应该都有同感:单点预测模型用到后面,瓶颈特别明显。仓库A的缺货可能是因为仓库B在集中调拨,爆款SKU的销量波动会沿着品类替代关系传导到周边商品——这些关联信息在传统时间序列模型里基本是浪费掉的。这两年我把图神经网络引入库存预测的落地实践中,确实解决了不少用LSTM、Prophet很难啃的场景。这个思路不是拿图模型替代所有统计模型,而是在真正有“网络结构”的库存体系里,把关联信息用起来。
我打算分几块把思路整理清楚:先拆解为什么库存预测需要图神经网络,再讲图结构怎么构建、模型怎么选、代码怎么跑通,最后是复盘实际落地时踩过的坑。整个过程会结合一个典型的多仓多SKU库存场景来讲,方便直接迁移到自己的业务里。
1. 库存预测的老问题:时序模型为什么越做越吃力
1.1 单点模型的天花板
标准的做法是每个SKU、每个仓库单独建模,序列进去,预测出来。LSTM、Transformer、Prophet都是这个套路。单点模型在数据平稳、波动小的时候表现还行,一旦遇到促销联动、仓间调拨、新品替代这些情况,预测误差会明显上升。原因不复杂:信息被限制了。每个点的模型只能看到自己的历史,看不到其他点上的变化。
举个实际案例。某快消品牌两个临近仓库,A仓负责核心城区,B仓负责郊区。消费者在A仓缺货时,订单会自然流转到B仓。如果只按各自历史销售预测,A仓的缺货信号还没体现出来之前,B仓的需求已经上来了。传统模型都捕捉不到这种联动,结果就是A仓持续缺货、B仓积压库存。这类问题靠加特征很难根治,因为特征工程能覆盖的关联是有限的,而且跨仓、跨品类的关联常常是非线性的。
1.2 库存系统的本质是网络
如果我们退一步看整个库存系统,你会发现它天然是图结构。节点是仓库、SKU或者两者的组合,边是它们的依赖关系。调拨关系、替代关系、共享同一批供应商、面向同一类客户群体,这些都能构成边。只要把这些关系显式建模出来,模型就有机会学到“联动”和“传导”。
这也是图神经网络(Graph Neural Network,GNN)入场的基本逻辑。GNN的核心能力是消息传递(message passing):每个节点聚合邻居的信息来更新自己。这个概念用在库存预测里非常自然——预测A仓某个SKU的销量,不只看它自己的历史,还把关联仓库、关联品类的状态一起聚合进来,相当于给预测模型装了“全局视野”。
2. 从业务问题到图建模:AI架构师的第一课
2.1 节点和边到底怎么定义
很多刚接触图神经网络的人上来就问“用什么模型”,其实第一步应该是“图怎么建”。图建错了,后面全是错。结合落地经验,节点定义通常有三种做法:
第一种,以仓库为节点。适合仓间调拨频繁、区域库存联动明显的场景。边表示调拨关系、共仓容约束或地理邻近性。这种建模粒度粗,数据量小,适合冷启动阶段验证方案。
第二种,以SKU为节点。适合品类替代、捆绑销售明显的场景。比如洗发水和护发素、手机和手机壳,销量之间存在强关联。边可以基于销量相关性、同品类关系或共同出现在一张订单里的频率来构建。
第三种,以仓库+SKU组合为节点。这是最细粒度、最贴近业务的建模方式。每个节点就是一个具体的“仓-品”实例,特征可以同时刻画仓库属性、SKU属性以及两者的交互。缺点是图规模大,计算成本高,一般需要采样和分布式训练支持。
边的权重也值得仔细斟酌。最简单的做法是0/1二值化,有关系为1,没关系为0。更精细的做法是用业务指标量化边强度,比如调拨频率、替代率、订单共现次数。我自己的经验是,边权不能拍脑袋设,最好从业务数据里统计出来,否则模型学到的是假关联。
2.2 特征工程:图的“数值状态”
有了图结构,接下来要解决节点特征的问题。节点特征相当于传统模型里的自变量,但这里要考虑把时序信息编码进去。
一个常用的方案是滑动窗口统计。对每个预测日,取过去28天的日销量、7天均值、28天标准差、缺货天数、促销标记等,构成一个固定长度的特征向量。对于仓库节点,还可以加上库容、覆盖区域人口密度等静态属性。对于SKU节点,可以加价格带、品类编码的embedding、生命周期阶段等。
这个阶段的经验是要注意特征的时间对齐。预测T+7的销量,特征只能用T日及之前的信息,绝对不能混入未来数据,否则回测结果虚高,上线就翻车。
2.3 标签设计跟上线节奏匹配
标签就是我们要预测的目标。多数场景是预测未来7天、14天或30天的累计销量。这里有个细节:如果业务上补货周期是7天,标签就对齐7天需求量;如果补货周期是14天,标签就取14天。标签和补货节奏错配,模型再准也无法直接辅助决策。
另外还要考虑延迟需求。缺货期间被压制的那部分需求,可能在补货后集中释放。这种延后效应如果没在标签里体现出趋势,模型的误差会在缺货修复期被拉大。比较稳妥的做法是早期先不考虑延迟需求,用纯销量做标签,等模型跑稳定后再逐步引入修正。
3. 图神经网络模型选型与原理浅讲
3.1 三种主流GNN架构的适用边界
图建好、特征做好后,才轮到模型选型。当前主流的GNN方案里有三种在库存预测场景中比较常用:图卷积网络(GCN)、GraphSAGE和图注意力网络(GAT)。
GCN是最基础的方案。它在频域定义卷积,通过拉普拉斯矩阵实现邻居信息聚合。优点是实现简单、计算快,适合图结构相对稳定、邻居数量均匀的场景。缺点是每一层的聚合权重是全局共享的,学不到“哪个邻居更重要”。
GraphSAGE的思路是采样邻居并做聚合,适合图规模大、没法一次性全量计算的场景。它支持均值聚合、LSTM聚合、池化聚合等多种方式,灵活性更高。在仓库数量几千、SKU数量几万的场景下,GraphSAGE基本是默认起点。
GAT引入了注意力机制,每个节点学习邻居的权重。聚合信息时,重要的邻居多分一些权重。在库存场景中,这意味着调拨频繁的仓库能拿到更大的注意力权重,而联系不强的节点则被自然忽略。GAT适合关系强度差异大的网络,但训练更慢,对图结构变化也更敏感。
从稳定性角度看,工业落地我更推荐先试GraphSAGE,因为它对图动态变化的容忍度最高。GCN对图结构突变比较敏感,GAT在边定义不准确时容易被错误注意力带偏。
3.2 消息传递机制与库存预测的结合点
GNN的核心是消息传递机制,本质上就是节点之间不断交换信息并更新自身的表征。一个典型的图神经网络层,可以拆成三个步骤:
第一步,每个节点收集邻居节点的特征信息(消息来源)。第二步,把这些信息按照一定规则聚合成一个向量(聚合函数,比如求和、均值或最大池化)。第三步,把聚合后的邻居信息与自己的特征拼接或加权融合,再经过一个非线性变换,得到更新后的节点表示(更新函数)。
把这个机制放到库存预测场景里理解。预测A仓SKU X在下一个周期的销量,消息传递过程是这样的:模型先把A仓SKU X历史销量特征表示出来,同时把与它有调拨关系的B仓、C仓的特征,以及与之有替代关系的SKU Y的特征一并收集过来。聚合之后,模型得到的不只是A仓SKU X自身的历史规律,而是整个局部库存子网络的动态态势。
多层的消息传递还能捕捉更远距离的关联。比如A仓的变化会影响B仓,B仓的变化又会影响C仓,两层消息传递之后,A仓的信息就能经过B仓传到C仓的表示里。这在供应链里对应的是间接联动,也就是那些看上去没有直接关系、但实际会通过中间节点互相影响的仓品组合。
这里有一个关键细节:邻居信息并不是简单平均一下就能用。不同邻居的影响力差异很大,尤其在库存网络里。调拨量大的邻居仓库,它的库存状态对你需求的挤压和补充作用,远大于那些只是地理上邻近但业务往来很少的仓库。这也是GAT这类注意力机制在理论上更贴合库存预测的原因。
3.3 与LSTM结合:时空建模的常用套路
纯GNN处理纯图信息没问题,但库存预测的核心输入还是时间序列。实际操作中,通常会把GNN和时间序列模型结合,形成时空预测架构。比较常用的套路是:先用LSTM或TCN对每个节点做时间编码,得到包含时序特征的节点表示;再把这些表示输入GNN层做空间信息聚合;最后接一个全连接层输出预测值。
这个架构在实践中比较稳定。以我自己跑过的项目为例,输入是每个节点过去28天的特征序列,LSTM编码后得到每个节点的隐状态,再经过两层GraphSAGE聚合邻居信息,最终输出未来7天的预测值。单从数值上看,在包含300个仓库节点和2000个SKU节点的数据集上,这种方法比单节点LSTM的预测误差(MAPE)降低约11%到17%,在促销期和平销期都有改善。
训练时要注意LSTM的序列长度和GNN的邻居数量需要调平衡。序列长度过长、邻居数量过大,训练时间会成倍增长,但收益会边际递减。一般建议序列长度取28到56天,邻居采样数量控制在10到20之间。
4. 手把手实现一个库存预测GNN模型
4.1 数据准备和图结构构建
为了演示,我构建一个模拟的多仓多SKU库存数据集。假设我们有5个仓库,50个SKU,生成365天的销售数据,包含季节性、趋势和随机波动。同时设定两个业务关系:仓库间调拨关系和SKU间替代关系。
用Python生成模拟数据,核心代码如下:
import numpy as np import pandas as pd from datetime import datetime, timedelta np.random.seed(42) warehouses = ['WH_A', 'WH_B', 'WH_C', 'WH_D', 'WH_E'] skus = [f'SKU_{i:02d}' for i in range(1, 51)] date_list = [] wh_list = [] sku_list = [] sales_list = [] start_date = datetime(2023, 1, 1) for day_offset in range(365): current_date = start_date + timedelta(days=day_offset) for wh in warehouses: for sku in skus: # 模拟季节性和趋势 seasonal = 1 + 0.3 * np.sin(2 * np.pi * day_offset / 365) trend = 1 + day_offset * 0.0005 # 模拟SKU和仓库的随机波动 sku_factor = 0.5 + np.random.rand() wh_factor = 0.7 + np.random.rand() * 0.6 base_sales = 10 * seasonal * trend * sku_factor * wh_factor sales = np.random.poisson(base_sales) date_list.append(current_date) wh_list.append(wh) sku_list.append(sku) sales_list.append(sales) df = pd.DataFrame({ 'date': date_list, 'warehouse': wh_list, 'sku': sku_list, 'sales': sales_list })这个模拟数据把销量拆成了季节因子、趋势因子、SKU因子和仓库因子的乘积,并用泊松分布加入随机噪声。结构上比较接近真实业务中的销量形态。
接下来构建图结构。仓库间的边基于调拨频率定义,SKU间的边基于历史销售相关性定义。
import networkx as nx from scipy.stats import pearsonr # 构建仓库调拨关系图 wh_graph = nx.Graph() wh_graph.add_nodes_from(warehouses) # 模拟调拨关系:相邻仓库之间有高概率存在调拨 transfer_pairs = [('WH_A', 'WH_B'), ('WH_B', 'WH_C'), ('WH_C', 'WH_D'), ('WH_D', 'WH_E'), ('WH_A', 'WH_C'), ('WH_B', 'WH_D')] for u, v in transfer_pairs: wh_graph.add_edge(u, v, weight=np.random.rand()) # 构建SKU替代关系图 sku_graph = nx.Graph() sku_graph.add_nodes_from(skus) # 基于销量相关性构建边 sales_pivot = df.pivot_table(index='date', columns='sku', values='sales', aggfunc='sum') corr_matrix = sales_pivot.corr() for i, sku_a in enumerate(skus): for j, sku_b in enumerate(skus): if i < j: corr_value = corr_matrix.loc[sku_a, sku_b] if abs(corr_value) > 0.6: sku_graph.add_edge(sku_a, sku_b, weight=corr_value)这里的边构建逻辑有两点需要注意。第一,仓库调拨关系应该来自ERP系统的调拨单,而不是随机生成。这里用随机数只是为了演示代码流程。第二,SKU替代关系的相关性阈值(0.6)需要根据实际数据分布调整,阈值设得太低会把无关SKU连在一起,太高则图太稀疏。
为了简化代码演示,这里先用同构图分别建模。实际业务中更推荐构建异构图,让仓库和SKU作为不同类型的节点,在一个图中统一表达。
4.2 特征工程与数据划分
节点特征用滑动窗口生成。窗口取28天,产出包括均值、标准差、最大值、最小值、趋势斜率等统计量。
def build_node_features(df, window=28): features = [] labels = [] node_ids = [] for (wh, sku), group in df.groupby(['warehouse', 'sku']): group = group.sort_values('date') sales_values = group['sales'].values for i in range(window, len(sales_values) - 6): hist = sales_values[i-window:i] feat = [ np.mean(hist), np.std(hist), np.max(hist), np.min(hist), hist[-1], np.polyfit(range(window), hist, 1)[0], np.mean(hist[-7:]), np.sum(hist[-7:]) ] label = np.sum(sales_values[i:i+7]) # 预测未来7天总销量 features.append(feat) labels.append(label) node_ids.append(f'{wh}_{sku}') return np.array(features), np.array(labels), node_ids数据划分要特别注意时序泄漏问题。不能用随机划分,要按时间顺序切分,比如前80%时间做训练集、后20%做测试集。
4.3 GNN模型搭建(基于PyTorch Geometric)
这里用PyTorch Geometric实现一个结合LSTM的GraphSAGE模型。先安装依赖:
pip install torch torch-geometric模型代码如下:
import torch import torch.nn as nn import torch.nn.functional as F from torch_geometric.nn import SAGEConv class LSTMGNN(nn.Module): def __init__(self, input_dim, hidden_dim, output_dim, seq_len, num_layers=2): super(LSTMGNN, self).__init__() self.seq_len = seq_len self.lstm = nn.LSTM(input_dim, hidden_dim, num_layers, batch_first=True) self.conv1 = SAGEConv(hidden_dim, hidden_dim) self.conv2 = SAGEConv(hidden_dim, hidden_dim) self.fc = nn.Linear(hidden_dim, output_dim) self.dropout = nn.Dropout(0.3) def forward(self, x, edge_index): # x: [num_nodes, seq_len, input_dim] batch_size, num_nodes, seq_len, input_dim = x.shape x = x.view(batch_size * num_nodes, seq_len, input_dim) lstm_out, _ = self.lstm(x) # 取最后一个时间步的输出 x = lstm_out[:, -1, :].view(batch_size, num_nodes, -1) # GNN 层聚合邻居信息 for conv in [self.conv1, self.conv2]: x = F.relu(conv(x, edge_index)) x = self.dropout(x) # 输出未来预测 x = self.fc(x) return x.squeeze(-1)这里的关键是数据形状的适配。PyTorch Geometric的SAGEConv接收的节点特征形状是[num_nodes, hidden_dim],但LSTM的输入需要[batch_size, seq_len, input_dim]。所以我在进入GNN层之前做了一个重构,把batch维和节点维合并,GNN层结束后再拆开。
4.4 训练、评估与调参要点
训练循环和普通PyTorch模型类似,这里不贴完整代码了,重点说几个实际跑起来容易踩坑的地方。
第一个是学习率。GNN模型在库存数据上的最优学习率通常在0.001到0.005之间,比纯LSTM模型的学习率稍低。原因是GNN层叠加之后梯度传播路径更长,学习率太大会导致训练震荡。我习惯用AdamW优化器,配合余弦退火学习率调度。
第二个是损失函数的选择。库存预测的标签是未来的销量,分布通常右偏,直接用MSE会让模型过分关注高销量节点。可以考虑对标签做log1p变换之后再算损失,或者用Huber Loss。log1p变换的做法适合销量跨度大的场景。
第三个是邻居采样。如果图规模太大,全量邻居聚合会导致显存溢出。PyTorch Geometric的NeighborSampler支持对每个节点的邻居做随机采样,一般采样10到20个邻居就足够了,再增加收益不大。
训练结束后,评估指标除了常见的MAPE、RMSE,还建议加一个业务指标:预测准确率(预测值落在真实值±20%区间内的比例)。这个指标更贴近实际决策场景,因为补货量一般有最小起订量和运输批次限制,微小的预测偏差不影响最终决策。
5. 上线部署与效果对比实录
5.1 项目效果:GNN与LSTM、Prophet的横向对比
为了验证方案有效性,我在一个真实的快消品库存数据集上做了横向对比。数据集包含约500个仓品组合节点,历史数据24个月,预测目标为未来7天销量。
| 模型 | MAPE(平销期) | MAPE(促销期) | 预测准确率(±20%) |
|---|---|---|---|
| Prophet | 32.4% | 41.7% | 51.2% |
| LSTM(单点) | 26.8% | 38.5% | 58.4% |
| GraphSAGE+LSTM | 23.1% | 31.2% | 66.9% |
整体来看,GNN方案在平销期比单点LSTM提升约14%,在促销期提升约19%。促销期的提升更明显,原因是促销期的销量联动效应更强,单点模型完全无法捕捉,而GNN可以利用网络结构把关联信息传递过来。
5.2 部署架构与推理性能
上线部署时,GNN模型面临的最大挑战是推理链路的延迟。库存预测一般是离线跑,频次是每天一次,对延迟不敏感。但如果做到实时补货建议,推理需要在秒级完成,这时就要做图分区和采样推理。
我采用的方案是:离线训练阶段做全量图计算,在线推理阶段对每个节点做局部采样。具体来说,预测某个仓品节点的销量时,只采样它的一阶和二阶邻居,然后做forward。这样单节点的推理耗时能控制在50毫秒以内,支持实时调用。
部署架构上,模型服务用TorchServe或ONNX Runtime,图结构数据通过Redis缓存,特征计算用Flink实时流水线完成。每天的预测任务跑在离线调度器上,结果写入数据仓库,下游补货系统直接读取预测表做补货建议。
5.3 模型监控与定期重训练
GNN模型上线后,监控不能只盯着预测误差。图结构本身可能会变化,比如新开仓库、淘汰SKU、调整调拨关系,这些都会影响模型的输入结构。
我的经验是:至少每周重新构建一次图,每月用全量数据重新训练一次模型。日常监控关注三个指标:预测准确率趋势、图结构变化频率、节点特征分布漂移。只要其中任何一个指标出现异常波动,就要触发告警并人工介入。
模型版本管理也要注意。用MLflow或类似的工具记录每次训练的数据版本、图结构版本、模型参数和评估指标。这样出了问题才能回溯到底哪一步导致预测质量下降。
6. 常见问题与踩坑实录
6.1 图构建不当导致预测效果反而变差
这是最容易踩的坑。很多人以为边的数量越多越好,实际不是。我在一个项目中把所有相关性大于0.3的SKU全部连边,结果图变得非常稠密,模型把大量噪声关联也学进去了,预测效果比单点LSTM还差。
后来把相关性阈值提高到0.7,同时只保留业务上确实存在替代关系的连接(比如同品牌同品类),效果才恢复正常。核心原则是边必须反映真实业务关系,相关性只是辅助判断信号,不能作为唯一依据。
6.2 冷启动的节点预测困难
新SKU和新仓库没有历史数据,节点的特征几乎为空,GNN的邻居聚合也无从谈起。这种情况需要在图结构上做一些特殊处理。
一种做法是把新节点连接到同类节点的代表节点上,比如同品类销量Top10的老SKU。另一种做法是退回到单点模型,用简单的统计方法先跑一段时间,等积累足够数据后再接入GNN。
我的建议是,在系统设计阶段就要预留冷启动通道。不要把GNN作为唯一的预测引擎,而是和其他方法混合使用,根据每个节点的数据量动态选择模型。
6.3 训练和推理的图结构不一致
这个问题比较隐蔽。训练时,整个图是完整的、静态的。到了推理时,实时数据会引入临时增加或断开的边。如果推理时用动态图结构与训练时的静态图结构差异太大,模型表现会明显下降。
解决办法是在训练阶段加入图结构扰动。具体做法是以一定概率随机删除或增加边,让模型对图结构的变化更鲁棒。这个操作类似图像领域的随机裁剪,效果很直接。
6.4 计算资源的分配策略
GNN的训练比传统时序模型重得多,尤其是图规模大的时候。GPU显存和训练时间都要考虑进去。
我的建议是分阶段控制成本。探索初期用小区块(比如一个区域的仓库)快速验证效果;确认可行后再扩展到全量数据。全量训练时,如果图太大,可以用Mini-batch训练配合邻居采样,而不是全图计算。
6.5 业务方的解释性要求
落地时经常会遇到业务方问:为什么这个SKU预测涨了30%?GNN本身的解释性比较弱,不像线性模型可以直接看系数。
我的处理办法是,在预测结果旁边附加“主要贡献节点”说明。具体做法是统计GNN注意力权重,找出对目标节点影响最大的Top5邻居,生成类似“WH_B仓该SKU近7天销量上涨20%,根据历史调拨关系,预计对WH_A仓需求产生传导影响”这样的解释文本。虽然不是模型内部机理的完全归因,但业务方接受度很高。
最后再分享一个小技巧
在项目推进的过程中,我逐渐意识到,GNN在库存预测上的价值,不单是预测精度提升了几个点,而是它让AI架构师重新理解了库存系统的结构。传统方法把所有仓品组合看成独立个体,GNN则把它们放回网络里,让模型能够感知到业务发生的真实环境。
如果你正准备在库存预测场景尝试GNN,我建议不要一上来就堆模型复杂度。先把图结构定义清楚,用最基础的GraphSAGE跑通基线,再逐步迭代。很多时候,业务关系的梳理和特征选择带来的收益,比换更复杂的GNN模型大得多。就拿我自己踩过的坑来说——有一版我换上了三层的GAT,效果反而不如两层的GraphSAGE,后来排查发现是图里有些边的权重设置不合理,注意力机制把错误的关系放大学习。反而是先把图构建逻辑理顺之后,模型效果一下子就上来了。
另外补一句经验,库存预测的工程链路比模型本身更考验功力。数据质量、特征时效性、图结构更新、模型监控,每个环节都可能成为瓶颈。GNN给了我们一个更强大的建模工具,但要把这个工具用出真正的业务价值,还是得回到对业务本身的理解上。