ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

指针网络与强化学习:求解TSP的完整PyTorch实战

2026/9/14 3:00:10 拓冰建站 浏览量
指针网络与强化学习:求解TSP的完整PyTorch实战 简介一套基于指针网络与强化学习求解旅行商问题TSP的Python代码实现主要面向对组合优化和深度强化学习感兴趣的开发者、研究者和竞赛学习者。项目采用简化训练策略不单独实现critic网络而是将最优路径长度作为critic值训练样本通过在[0,1]×[0,1]网格上均匀采样二维点生成并由Concorde求解器给出最优解作为监督信号整体思路清晰适合作为相关方向的入门实践。压缩包共14个文件大小约4.01MB其中包含8个Python脚本覆盖模型、训练器、数据加载、配置与工具函数等环节另有2张结果图片、2个测试数据文件、1份README说明文档和1个gitignore文件目录结构简洁便于按需阅读。已有866人学习/下载。通过这套代码读者可以复现TSP10与TSP50的强化学习训练流程查看经过100,000步训练后的测试结果并使用diff指标对比强化学习解与最优解的差距。对于想动手实践深度强化学习、路径规划或组合优化的读者这份实现提供了可运行的完整示例和清晰的代码组织方式更详细的依赖配置与使用方法可阅读包内README说明。1. 用指针网络给 TSP 找一条最短环游为什么值得自己做一遍先划定一个具体问题给定 n 个城市的二维坐标求一条访问每个城市恰好一次并回到起点的最短闭合路径这就是经典的 TSP旅行商问题。当 n 从 20 涨到 200精确求解器的耗时立刻变得不可接受工程上常用的做法是 LKH、OR-Tools 这类启发式搜索它们快但每个实例都要重新搜索。指针网络Pointer Networks换了一条路它学习的是给定坐标直接输出访问顺序的映射训练好之后一次前向传播就能给出一个可用解不需要在推理时做任何搜索。而配合强化学习训练模型只需要以路径长度作为奖励连最优解标签都不用准备——这对入门深度强化学习、想快速看清端到端组合优化这套玩法到底行不行的工程师来说是性价比很高的一个实验。本文会按模型结构、训练方法、完整 Python 实现、最后到验证与调参的顺序把整条链路摊开讲。2. 指针网络的结构把“输出词表”换成“输入下标”2.1 标准 seq2seq 为什么解不了 TSP词表是固定大小的做过机器翻译或文本生成的人对 seq2seq 都很熟编码器把输入句子编码成一系列隐状态解码器每一步通过注意力机制从这些隐状态里汇总一个上下文向量再经过一个线性层和 softmax从固定词表里挑一个词。这里的关键限制在于输出词表的大小是预先固定的。TSP 的输出是 n 个城市的某种排列城市数量一变输出维度就得跟着变更麻烦的是这 n 个词的含义就是输入序列本身。你当然可以把坐标喂给编码器、让解码器从一个大词表里选城市 ID但这个词表要么大到浪费要么在 n 变化时完全失效。指针网络的改动非常直接注意力打分函数不再用来做加权求和而是直接输出一个在输入位置上的概率分布然后取 argmax 或采样作为当前步的输出。换句话说模型不再说我选词表里的第 5 个词而是说我选输入序列里的第 5 个城市。这个设计让输出空间与输入序列长度自动对齐n 变了模型结构也不需要动。从信号处理的角度看它本质上是 content-based 的硬匹配比软性注意力更贴合从一组元素中挑一个出来这类组合问题。2.2 最小可运行的 Pointer Network 实现PyTorch指针网络的核心组件是编码器、解码器和打分函数。下面这份代码可以直接保存成 model.py包含了完整的模型定义和前向逻辑。import torch import torch.nn as nn import torch.nn.functional as F class PointerNetwork(nn.Module): def __init__(self, embed_dim128, hidden_dim256): super().__init__() self.encoder_embed nn.Linear(2, embed_dim) # 坐标 - 向量 self.encoder nn.LSTM(embed_dim, hidden_dim, batch_firstTrue) self.decoder_embed nn.Linear(2, embed_dim) self.decoder nn.LSTM(embed_dim, hidden_dim, batch_firstTrue) # 注意力打分的三组投影参数 self.W1 nn.Linear(hidden_dim, hidden_dim, biasFalse) self.W2 nn.Linear(hidden_dim, hidden_dim, biasFalse) self.v nn.Linear(hidden_dim, 1, biasFalse) def forward(self, coords, decode_typegreedy): # coords: [B, n, 2]欧氏坐标范围 [0,1] B, n, _ coords.shape enc_emb torch.relu(self.encoder_embed(coords)) enc_outs, (h, c) self.encoder(enc_emb) # enc_outs: [B, n, hidden] # 固定城市 0 为起点解码器第一步从剩余城市里选 inp torch.relu(self.decoder_embed(coords[:, 0:1, :])) mask torch.zeros(B, n, devicecoords.device) mask[:, 0] 1 outputs, log_probs [], [] for step in range(n - 1): _, (h, c) self.decoder(inp, (h, c)) dec_state h[-1] # [B, hidden] # 对每个输入位置打分取概率最大的下标作为输出 scores self.v(torch.tanh( self.W1(enc_outs) self.W2(dec_state).unsqueeze(1) )).squeeze(-1) # [B, n] scores scores.masked_fill(mask.bool(), -float(inf)) probs F.softmax(scores, dim-1) if decode_type greedy: idx probs.argmax(dim-1) logp None else: # sampling 训练模式 dist torch.distributions.Categorical(probs) idx dist.sample() logp dist.log_prob(idx) mask mask.scatter(1, idx.unsqueeze(1), 1.0) outputs.append(idx) log_probs.append(logp) # 把当前选中的城市坐标作为解码器下一步输入 picked torch.gather(coords, 1, idx.unsqueeze(-1).expand(-1, -1, 2)) inp torch.relu(self.decoder_embed(picked)) pi torch.stack(outputs, dim1) # [B, n-1]不含起点 return pi, log_probs模型前向的流程可以拆成四步看。第一步把 [B, n, 2] 的坐标经过一个线性层映射成 embedding再过 LSTM 得到每个城市的编码隐状态。第二步解码器每一步把上一步选中的城市坐标作为输入推进 LSTM 得到当前隐状态。第三步用 W1 和 W2 分别投影编码器输出与解码器状态相加后过 tanh再用 v 压缩成标量这个标量就是该城市被选中的分数。第四步softmax 后直接当作选择概率贪婪解码时取 argmax训练时用 Categorical 采样。这里有几个参数值得注意。embed_dim 是坐标嵌入的维度一般 128 足够n 超过 100 时可以升到 256hidden_dim 是 LSTM 隐层宽度直接决定了打分函数的表达能力但也不是越大越好hidden_dim 增加会让 W1 和 W2 的矩阵变大注意力计算量随 n 线性增长显存压力来自这里。mask 的处理方式是把已访问城市的分数置为负无穷而不是把概率置零这一点后面会单独讲。2.3 别忘了 mask这是“每个城市只访问一次”的唯一保障如果只看 2.2 的代码最容易被忽略的就是 mask 那一行。TSP 的约束是每个城市恰好访问一次这意味着模型每一步的输出都不能是已经选过的城市。最简单的做法是维护一个 [B, n] 的 0/1 掩码选中某城市后通过 scatter 把对应位置置 1。关键点在于 mask 必须作用在 softmax 之前的 scores 上用-float(inf)把这些位置屏蔽掉。如果先算 softmax 再从概率里强制置 0这一项在反向传播时仍然会有梯度残留模型会收到这个城市虽然不该选但它的分数也在被优化的错误信号训练很难收敛。我在实际调试中还遇到过一种隐蔽的错误mask 没有包含起点城市导致解码器第一步就把起点再选一次。这在训练初期 loss 看起来正常但最终输出一定是错的。所以代码里mask[:, 0] 1这行虽然简单却是整条数据流正确的关键前提。2.4 静态欧式 TSP 的数据约定与训练/推理模式这套实现针对的是静态欧式 TSP即每个实例的城市坐标一次性给出距离按欧氏距离计算且实例之间相互独立。与此相对的是动态 TSP城市会随智能体移动而更新那种情况需要模型每步重新编码指针网络的原版结构就不够用了。坐标生成通常用 [0,1] 区间上的均匀分布这是组合优化领域最常见的 benchmark 约定路径长度的数值范围也合适n20 时最短环游通常在 4 左右。如果坐标范围改成 [0, sqrt(n)]距离量级会变大梯度尺度也随之变化一般建议统一用 [0,1]。这里还有一个容易踩的坑训练集和测试集应该各自在线生成而不是事先固定一个数据集。原因在于强化学习的目标是泛化到任意随机实例上提前固定数据集会让模型记住局部模式。训练时每个 batch 用torch.rand(batch_size, n, 2)重新生成坐标相当于无限训练集。推理时用 greedy 解码这和训练时的 sampling 解码要区分开后面会看到rollout baseline 用的正是 greedy 解码作为策略评估方式。3. 策略梯度训练REINFORCE 与 rollout baseline3.1 先想清楚为什么不能直接用交叉熵训练如果给指针网络配上最优路径标签完全可以按监督学习训练用交叉熵让模型输出逼近标签。但这么做有两个问题一是要先用求解器生成最优标签n 稍大标签成本就很高二是监督学习学的是模仿求解器的输出而不是直接优化路径长度求解器的偏好会被模型照单全收。强化学习的思路完全不同把路径长度作为奖励信号模型自己探索出比标签更好的解也是允许的。先把问题写成 MDP。状态是当前已访问城市集合和当前所在城市动作是选择下一个未访问城市转移就是把该城市加入已访问集合奖励在整条轨迹结束时给出等于环游总距离的负值。策略就是解码器每一步输出的选择概率。轨迹长度固定为 n-1 步每一步的合法动作由 mask 限定这个设定让问题变成了一个标准的有限步序贯决策问题。3.2 REINFORCE with baseline方差要降符号不能反直接优化期望路径长度 J(θ) E[L(π)]对 θ 求梯度会得到策略梯度。由于 L 是离散排列的确定性函数不可微梯度只能通过 log 概率来估计。在最小化路径长度的设定下策略梯度的更新方向可以用下面这个式子理解∇J(θ) E[ (L(π) - b) · ∇θ log p(π) ]这里的 b 是 baseline它的作用是降低方差不改变梯度的期望值。选择 b 的常见做法有两种。一种是训练一个 critic 网络回归当前状态下的期望路径长度另一种是 rollout baseline即维护一个稍旧版本的策略用它的 greedy 解码结果作为 baseline。对于 n20 到 n50 这个规模我建议直接用 rollout baseline少一个 critic 网络就少一套调参负担而且大规模实验比如 Kool 等人后来的 Attention Model已经验证过 rollout baseline 在小规模 TSP 上稳定且高效。一个特别容易出错的细节是符号方向。上面公式是最小化期望长度的写法所以 loss 写成log_probs * (L - b)。当采样解比 baseline 好L 更小时L - b 为负loss.backward() 会增大对应路径的 log 概率这才是正确的鼓励方向。如果写成b - L那就是在反向惩罚好解训练时你会看到平均环长不降反升。3.3 训练循环代码与核心参数训练循环需要两个辅助函数计算环游总长度以及单步训练。下面是完整的实现可直接放到 train.py 中。def compute_tour_length(coords, pi): # pi: [B, n-1]不含起点把起点 0 拼到头尾后计算欧氏环长 B pi.size(0) zeros torch.zeros(B, 1, dtypetorch.long, devicepi.device) full torch.cat([zeros, pi], dim1) # [B, n] nxt torch.cat([full[:, 1:], full[:, :1]], dim1) # 回到起点 xy1 coords.gather(1, full.unsqueeze(-1).expand_as(coords)) xy2 coords.gather(1, nxt.unsqueeze(-1).expand_as(coords)) return (xy1 - xy2).pow(2).sum(-1).sqrt().sum(-1) # [B]compute_tour_length 先把起点 0 拼到头部再把 full 序列整体左移一位作为下一跳目标这样构造出的首尾相接序列就是一条完整环游。gather 操作按行取出坐标逐段计算欧氏距离后求和。需要注意 full 和 nxt 都是 [B, n] 的索引矩阵full.unsqueeze(-1).expand_as(coords)会把索引广播成与 coords 相同的 [B, n, 2] 形状gather 才能逐维对齐。def train_step(model, optimizer, batch_size512, n_cities20): coords torch.rand(batch_size, n_cities, 2, devicedevice) pi, log_probs model(coords, decode_typesampling) log_probs torch.stack([lp for lp in log_probs], dim1) # [B, n-1] length compute_tour_length(coords, pi) with torch.no_grad(): pi_base, _ model(coords, decode_typegreedy) len_base compute_tour_length(coords, pi_base) advantage (length - len_base).detach() loss (log_probs * advantage.unsqueeze(1)).mean() optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() return loss.item(), length.mean().item(), len_base.mean().item()训练循环里值得关注的点有三个。第一baseline 的计算放在torch.no_grad()里greedy 解码本身不需要梯度也不应该把 baseline 的梯度混入主loss。第二advantage用了detach()因为 advantage 只作为加权系数不需要对长度求导。第三梯度裁剪用clip_grad_norm_限制在 1.0这能防止采样中偶发的极端路径造成梯度爆炸是强化学习训练里几乎必备的一行。下面是这份实现里最常用的参数组合可以当作起点再根据收敛曲线调整。参数n20 建议值n50 建议值说明embed_dim128128坐标嵌入维度n 更大时升到 256hidden_dim256256LSTM 隐层宽度batch_size512256大 batch 降低策略梯度方差learning_rate1e-31e-3Adam 默认 lr过大容易震荡max_norm1.01.0梯度裁剪阈值训练步数3000~50008000~15000观察 avg_len 曲线决定batch_size 对强化学习的影响比监督学习更明显。策略梯度的方差来自采样batch 越大advantage 均值的方差越小训练就越稳定。但显存有限时优先降 batch 而不是降 hidden_dim因为 batch 影响的是梯度估计质量hidden_dim 影响的是模型表达能力两者不可互相替代。4. 从零跑通 Python 实现训练、可视化与排错4.1 工程目录与依赖准备与其到处找源码包下载不如按下面的文件结构把工程搭起来每个文件职责单一调试时也更容易定位问题。pointer_nets_tsp/ ├── model.py # PointerNetwork 类与 compute_tour_length ├── train.py # 训练循环、日志打印 ├── eval.py # 生成实例、可视化环游、OR-Tools 对比依赖只有五个Python 3.10 以上、PyTorch 2.0 以上、numpy、matplotlib、tqdm。验证 gap 时需要额外安装 ortools。CPU 也能跑通 n20 的最小验证只是 batch 建议降到 128训练步数加到 8000 左右如果机器上有 NVIDIA GPU一个 4G 显存的卡就能让 n50 的整个训练跑得很舒服。4.2 训练入口与日志监控train.py 的入口部分很短核心逻辑就是循环调用 train_step。这里的关键是固定随机种子否则强化学习训练结果很难复现。import torch from model import PointerNetwork from train import train_step device torch.device(cuda if torch.cuda.is_available() else cpu) torch.manual_seed(42) model PointerNetwork().to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-3) for step in range(5000): loss, avg_len, base_len train_step(model, optimizer) if step % 200 0: print(fstep {step:5d} | loss {loss:.3f} | favg_len {avg_len:.3f} | base {base_len:.3f})日志打印的三列分别代表策略梯度 loss、当前策略采样解的平均环长、baseline 的 greedy 环长。n20 时avg_len 会在大约 2000 步内从初始的 7~8 降到 5 以下最终逼近 4.2 左右如果训练 5000 步 avg_len 还在 6 以上优先检查 mask 是否正确、l r 是否过大而不是加大模型。4.3 可视化画一条环游出来看交叉训练到中途就可以随机生成一个实例用 greedy 解码画图检查这一步对判断模型是否学到了空间结构至关重要。import matplotlib.pyplot as plt def plot_tour(coords, pi): tour [0] pi [0] # 补上起点并回到起点 x, y coords[tour, 0], coords[tour, 1] plt.figure(figsize(6, 6)) plt.plot(x, y, o-, linewidth1.5) plt.title(fn{len(coords)}, length{len(coords):.3f}) plt.show()画出来的环游如果存在明显的交叉线段说明模型还在局部最优附近挣扎一旦交叉消除、路径变得平滑基本可以判断训练到位。可视化能和 avg_len 曲线互相印证避免只看数字做出错误判断。4.4 常见失败模式速查表现象常见原因处理方法loss 在降avg_len 不动模型退化为固定偏序输出调小 lr调大 batch_size输出序列出现重复城市mask 没作用在 scores 上或漏掉起点在 softmax 前 masked_fillloss 出现 NaN某一行 scores 全为 -inf检查 mask 是否覆盖所有城市训练震荡、avg_len 波动大advantage 方差过大batch 内标准化或提高 baseline 频率最终环长与最优解 gap 5%LSTM 表达力不足换 Transformer 编码器AM 方向或加 hidden_dim5. 三个让训练更稳的进阶技巧进阶一在 batch 内做 advantage 标准化。上面代码里length - len_base的量级随 n 不同差别很大n20 时方差小n50 时方差可能放大一个数量级导致同一份学习率在不同 n 上表现截然不同。常见的修法是一行标准化代码adv (length - len_base).detach() adv (adv - adv.mean()) / (adv.std() 1e-9) # 1e-9 防除零这样 advantage 变成了相对好坏的度量梯度尺度不再依赖环长的绝对值lr 在 n20 和 n50 时可以共用一组值。进阶二rollout baseline 不能一直不更新。训练初期策略快速变强baseline 如果冻结太久length - len_base持续为负且绝对值偏大虽然方向正确但方差会升高。更新的节奏很有讲究每 500 步在同一个固定验证集上比较当前策略和 baseline 的 greedy 环长如果当前策略平均短 2% 以上就把当前权重整体复制给 baseline再继续训练。更新太频繁会让 advantage 接近零、学习信号变弱太稀疏方差又太大2% 阈值是一个足够保守且稳定的经验值。进阶三用 OR-Tools 的参考解算 gap而不是自己感觉“看起来不错”。OR-Tools 的 RoutingModel 内置了 TSP 求解器少量代码就能得到高精度参考解def ortools_tour(coords, time_limit5): from ortools.constraint_solver import pywrapcp, routing_enums_pb2 n len(coords) manager pywrapcp.RoutingIndexManager(n, 1, 0) routing pywrapcp.RoutingModel(manager) def dist(i, j): return int(round(10000 * float(torch.linalg.vector_norm(coords[i] - coords[j])))) cb routing.RegisterTransitCallback(dist) routing.SetArcCostEvaluatorOfAllVehicles(cb) params pywrapcp.DefaultRoutingSearchParameters() params.time_limit.seconds time_limit sol routing.SolveWithParameters(params) route, idx [], routing.Start(0) while not routing.IsEnd(idx): route.append(manager.IndexToNode(idx)) idx sol.Value(routing.NextVar(idx)) return route注意距离回调要求返回整数乘 10000 是为了保留欧氏距离的精度。在 n20 上OR-Tools 默认 5 秒内给出的解已经非常接近全局最优把它作为参考值去算(ours - ref)/ref * 100%的 gap就是模型质量最客观的度量。n50 之后 LSTM 类指针网络的 gap 会明显上升这是模型容量和 attention 机制的边界问题如果需要更高的解质量可以把编码器替换成 Transformer 结构的 Attention Model那已经是这篇文章之后的下一个工程了。本文还有配套的精品资源点击获取