PyTorch Geometric 异构图实战:把供应链运输成本预测完整落地指南
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
速览:用 PyTorch Geometric(PyG)把供应商、仓库、客户、产品四类实体装进同一张异构图,对「仓库→客户」边做运输成本回归——从最小可跑版本、切分可信度、时序采样防泄漏,一路走到分布式采样与 torch.jit 交付,链路一次走全。适合有 Python/PyTorch 基础、没碰过图学习的工程师,读完可以直接照抄骨架去接自己的物流数据。
一、先跑通:PyG 异构图最小可跑版本
不等概念铺完,先让代码转起来。PyG 的HeteroData用「节点类型 + 边类型」组织数据,边统一是 2×E 的索引:第 0 行起点、第 1 行终点。
import torch from torch_geometric.data import HeteroData data = HeteroData() data['supplier'].x = torch.randn(120, 8) # 产能、区位、履约率(z-score) data['warehouse'].x = torch.randn(30, 8) # 库容、周转天数、租金 data['customer'].x = torch.randn(5000, 8) # 下单频次、账期、区域 data['product'].x = torch.randn(300, 8) # 体积重、温层、单价 data['supplier', 'supplies', 'warehouse'].edge_index = sup_wh # 2xN data['warehouse', 'stores', 'product'].edge_index = wh_prod # 2xN data['warehouse', 'transports', 'customer'].edge_index = wh_cust # 待预测边模型侧只需两段:SAGEConv编码器负责消息传递,解码器把边两端点向量拼起来过 MLP 输出标量。
from torch_geometric.nn import SAGEConv, to_hetero class Encoder(torch.nn.Module): def __init__(self, h, o): super().__init__() self.conv1 = SAGEConv((-1, -1), h) # -1 表示维度由数据推断 self.conv2 = SAGEConv((-1, -1), o) def forward(self, x, edge_index): return self.conv2(self.conv1(x, edge_index).relu()) class EdgeDecoder(torch.nn.Module): def __init__(self, h): super().__init__() self.mlp = torch.nn.Sequential( torch.nn.Linear(2 * h, h), torch.nn.ReLU(), torch.nn.Linear(h, 1)) def forward(self, z, edge_index): row, col = edge_index z = torch.cat([z['warehouse'][row], z['customer'][col]], -1) return self.mlp(z).view(-1) # 标量:单均运输成本 class CostModel(torch.nn.Module): def __init__(self, h): super().__init__() self.enc = to_hetero(Encoder(h, h), data.metadata(), aggr='sum') self.dec = EdgeDecoder(h) def forward(self, x_dict, e_dict, e_label): return self.dec(self.enc(x_dict, e_dict), e_label)import torch.nn.functional as F model = CostModel(64) opt = torch.optim.Adam(model.parameters(), lr=0.01) et = ('warehouse', 'transports', 'customer') pred = model(data.x_dict, data.edge_index_dict, data[et].edge_label_index) loss = F.mse_loss(pred, data[et].edge_label) # 回归,MSE 起步 loss.backward() opt.step()跑通之后再看两个「为什么」。其一,成本不是单边属性:供应商产能掉、某仓缺货、改走另一条线,影响是沿着边一跳跳传过去的,任何单表 SQL 都截不住这条链;图的价值就是让消息沿关系走。其二,边级预测的本质是「读两端猜整条边」——编码器把节点压成向量,解码器只消费拼接后的端点向量,所以模型规模可以和边数解耦。
二、把预测做可信:反向边防泄漏与指标口径
模型能跑不等于结果能用。先看最隐蔽的坑:训练/验证/测试要按「边」切,而不是按「节点」。
切边只切正向边、反向边原封不动留在原图,采样子图会把测试边带进训练批次——测试 RMSE 虚低,上线即失真。
RandomLinkSplit用rev_edge_types把反向边一起切走:
from torch_geometric.transforms import RandomLinkSplit train, val, test = RandomLinkSplit( num_val=0.1, num_test=0.1, neg_sampling_ratio=0.0, # 回归不需要负样本 edge_types=[et], # 待预测边 rev_edge_types=[('customer', 'rev_transports', 'warehouse')], )(data)评估侧同时算 RMSE 和 MAE,并且直接换算成钱,方便和业务对账:
@torch.no_grad() def evaluate(d): model.eval() p = model(d.x_dict, d.edge_index_dict, d[et].edge_label_index) y = d[et].edge_label.float() rmse = float(F.mse_loss(p, y).sqrt()) # 千元/单 mae = float(F.l1_loss(p, y)) return rmse, mae假设测试集 MAE 是 0.4 千元/单,月均单量 8 万,平均偏差约为 3.2 万元/月。拿这个数和现行固定报价方案的实际偏差比,值不值得上线一眼可见。
val 的数字只用于早停,对外汇报一律取 test split;拿 val 当成绩,等于用考卷原题备考。
best, bad = float('inf'), 0 for epoch in range(1, 201): train_one_epoch(train) # 省略:前向 + MSE + step rmse_v, _ = evaluate(val) best, bad = (min(best, rmse_v), bad + 1) if rmse_v < best else (best, bad + 1) if bad >= 8: break # 早停,8 轮无改善 print(evaluate(test)) # 只信这里三、时间进场:时序邻居采样如何防未来泄漏
运输关系每天都在变,上周才开的线路不该出现在训练样本里。静态切边到这里就到头了,需要换LinkNeighborLoader按时序采样。
from torch_geometric.loader import LinkNeighborLoader loader = LinkNeighborLoader( data, num_neighbors=[5, 5], edge_label_index=(et, wh_cust), # 待预测边 edge_label_time=edge_time - 1, # -1:只允许采样到过去 time_attr='time', temporal_strategy='last', # 每跳截断在预测时点之前 batch_size=256, shuffle=True, )temporal_strategy='last'配合edge_label_time意味着:每一跳采样都只能拿到该边「预测时点」之前已经存在的关系,未来边在机制上不可见,而不是靠事后清洗。这在物流数据里比模型结构更容易翻车——线上口径对不上,多半先查这里。
评估口径换成推荐式的:torch_geometric.metrics提供LinkPredPrecision(k)与LinkPredRecall(k)。Precision@20 读作「给每条线路推 20 个候选合作方,平均几个是真发生过往来的」;召回回答「真实合作被 Top-20 覆盖了多少」。转成链路预测任务时记得在 loader 里加neg_sampling=dict(mode='binary', amount=2)造负样本——负样本比例对 Precision@k 影响很大,要调参而不是写死。
四、上规模与交付:分布式采样和 torch.jit 导出
节点和边超出单机内存之后,PyG 的torch_geometric/distributed/分两级扩展。
先用Partitioner离线切图,产物是META.json、节点/边映射表加若干分片,每片带graph.pt与node_feats.pt:
from torch_geometric.distributed import Partitioner Partitioner(data, num_parts=4, root='./parts').partition() # 落盘:META.json + node_map/ + edge_map/ + part0..part3/训练侧换成DistNeighborLoader,绑定本分片后,本地邻居直接读盘,跨分片邻居走 RPC 异步拉取,采样范围从全图收窄到「本机分片 + 一跳远程」,吞吐随机器数近似线性扩——订单边动辄上亿时这一步基本是必选项。
交付侧,模型用torch.jit.script整体导出,推理环境不再依赖训练时的依赖树:
scripted = torch.jit.script(model) torch.jit.save(scripted, 'supply_chain.pt') loaded = torch.jit.load('supply_chain.pt') out = loaded(x_dict, edge_index_dict, edge_label_index)导出的对象是「编码器 + 解码器」,输入仍是x_dict/edge_index_dict,线上服务把特征拼好直接喂入即可;如果线上只迭代编码器,也可以单独导出编码器复用既有服务,省一次全量发布。
五、上线前检查清单
RandomLinkSplit的rev_edge_types是否与正向边一一对应,反向边没有留在训练集里- 时序 loader 的
edge_label_time是否减过 1,测试边不存在于任何训练批次的时间窗之后 - 对外只报 test split 的 RMSE/MAE,验证集数字仅作早停依据
Partitioner落盘后核对每个分片的graph.pt、node_feats.pt齐全,META.json与机器数一致torch.jit.load在干净环境跑通一次端到端推理,输入输出 shape 与训练日志一致neg_sampling比例已在离线网格上调过,而非沿用默认值
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考