ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

指针网络+强化学习求解TSP:从零搭建与避坑指南

指针网络+强化学习求解TSP:从零搭建与避坑指南 简介这份资源围绕指针网络求解旅行商问题TSP的强化学习实现展开面向具备一定Python与深度学习基础、希望动手复现组合优化算法的开发者与研究者。代码以最佳路径长度作为critic值省去额外critic网络训练样本由[0,1]×[0,1]网格均匀采样生成最优解借助Concorde求解需将其加入系统PATH。压缩包共14个文件约4.01MB包含8个py源码模型、训练器、数据加载、层定义与配置等、2个npz测试数据集、2张结果图、1份readme说明及gitignore结构清晰便于按模块阅读。资源中给出TSP10在10万步训练后的测试结果并以diff指标衡量强化学习解与最优解的差距读者可据此评估模型收敛与泛化表现。目前已有867人学习下载适合作为组合优化与强化学习交叉方向的入门实践参考。1. 指针网络做 TSP 强化学习为什么它比传统启发式更值得投入如果你做过 TSP旅行商问题求解大概率经历过这样的场景城市规模从 20 涨到 100原本跑得好好的遗传算法突然变得不稳定换 CPLEX 又嫌太重手写 2-opt 邻域搜索调参调到怀疑人生。指针网络Pointer Network配合强化学习恰好切入了这个痛点——它不需要标注好的最优解直接用奖励信号驱动模型学会输出城市访问序列推理时一次前向传播就能给出解速度比迭代式启发式快几个数量级。这个方向适合两类人一是想入门深度强化学习但不想碰游戏环境的 Python 开发者TSP 的奖励函数天然清晰没有稀疏奖励的玄学问题二是做组合优化落地的工程师需要一套能泛化到不同城市规模的求解框架。本文从零搭建一套可运行的指针网络 REINFORCE 训练流程覆盖数据生成、模型定义、训练循环、贪心与采样解码、避坑排查最后给出规模化验证的实用技巧。代码基于 PyTorchPython 3.8 以上即可跑通。2. 指针网络与 REINFORCE 的配合逻辑为什么不用交叉熵2.1 指针网络解决的是「输出字典随输入变化」的问题标准 Seq2Seq 模型在解码时输出词表是固定的。TSP 不一样输入 10 个城市输出就是 10 个位置的排列输入 50 个城市输出就是 50 个位置的排列。词表大小随输入变化固定 softmax 层没法处理。指针网络的核心改动是解码每一步不再从固定词表选 token而是用注意力机制计算当前解码状态与所有编码器隐藏状态的匹配分数归一化后作为指向输入位置的指针概率分布。数学上第 $t$ 步指向城市 $i$ 的概率为$$p(i \mid \text{context}) \text{softmax}(u_i)$$其中 $u_i v^T \tanh(W_1 e_i W_2 d_t)$$e_i$ 是编码器对城市 $i$ 的输出$d_t$ 是解码器当前状态。这个设计让模型天然支持变长输入且输出必然是输入的一个排列配合 mask 机制。2.2 为什么用 REINFORCE 而不是监督学习监督学习需要标注最优解而 TSP 最优解在 50 城市以上就极难获取。REINFORCE 属于策略梯度方法直接用路径长度的负值作为奖励$$\nabla_\theta J(\theta) \approx \frac{1}{B} \sum_{b1}^{B} (L(\tau_b) - b) \nabla_\theta \log p_\theta(\tau_b)$$其中 $L(\tau_b)$ 是第 $b$ 条采样路径的总长度$b$ 是基线baseline用于降低方差。常见做法是用贪心解码的路径长度作为基线这样不需要额外训练 Critic 网络实现简单且效果稳定。注意基线不参与梯度回传只做数值减法。如果用 Critic 网络做基线需要 detach 后再减。2.3 最小可运行代码数据生成与模型定义先解决数据。TSP 实例生成很简单在单位正方形内均匀采样 N 个点。import torch import torch.nn as nn import torch.nn.functional as F def generate_tsp_data(batch_size, num_cities, device): 生成 batch_size 个 TSP 实例每个实例 num_cities 个城市坐标 # 坐标范围 [0, 1]形状 (batch_size, num_cities, 2) return torch.rand(batch_size, num_cities, 2, devicedevice)指针网络模型分编码器和解码器两部分。编码器用 LSTM 或 Transformer 均可这里用 LSTM 做最小实现。class PointerNetwork(nn.Module): def __init__(self, input_dim2, hidden_dim128): super().__init__() self.hidden_dim hidden_dim # 编码器将城市坐标序列编码为隐藏状态 self.encoder nn.LSTM(input_dim, hidden_dim, batch_firstTrue) # 解码器输入是上一步选中的城市坐标 上一步隐藏状态 self.decoder nn.LSTMCell(input_dim, hidden_dim) # 注意力参数将编码器输出和解码器状态映射为指针分数 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, x, decode_typesampling): x: (batch, num_cities, 2) decode_type: sampling 用于训练greedy 用于基线和推理 batch_size, num_cities, _ x.shape # 编码 encoder_out, (h, c) self.encoder(x) # encoder_out: (B, N, H) # 解码器初始状态用编码器最后一步的隐藏状态 decoder_h h.squeeze(0) # (B, H) decoder_c c.squeeze(0) # 初始输入一个可学习的起始向量这里简化为全零 decoder_input torch.zeros(batch_size, 2, devicex.device) # 记录已访问城市防止重复选择 mask torch.zeros(batch_size, num_cities, devicex.device) # 记录路径 pointers [] for _ in range(num_cities): decoder_h, decoder_c self.decoder(decoder_input, (decoder_h, decoder_c)) # 注意力分数计算 query self.W2(decoder_h).unsqueeze(1) # (B, 1, H) keys self.W1(encoder_out) # (B, N, H) scores self.v(torch.tanh(query keys)).squeeze(-1) # (B, N) # 已访问城市分数置为极小值 scores scores.masked_fill(mask 1, -1e9) probs F.softmax(scores, dim-1) if decode_type greedy: idx probs.argmax(dim-1) else: idx torch.multinomial(probs, 1).squeeze(-1) pointers.append(idx) mask mask.scatter(1, idx.unsqueeze(1), 1) # 下一步输入是当前选中城市的坐标 decoder_input x[torch.arange(batch_size), idx] return torch.stack(pointers, dim1) # (B, N)这段代码里几个关键点mask保证每个城市只被访问一次decode_type控制训练时用采样、评估时用贪心注意力分数计算采用加性注意力比点积注意力更适合小规模 TSP。参数hidden_dim建议从 128 起步城市数超过 50 时加到 256。3. 训练循环与奖励设计从随机路径到稳定收敛3.1 奖励函数与损失计算TSP 的奖励就是路径长度的负值。给定指针序列计算总距离def compute_tour_length(x, pointers): x: (B, N, 2) 城市坐标 pointers: (B, N) 访问顺序 返回: (B,) 每条路径的总长度 batch_size, num_cities, _ x.shape # 按 pointers 重排城市坐标 ordered x[torch.arange(batch_size).unsqueeze(1), pointers] # (B, N, 2) # 计算相邻城市距离包括首尾闭合 diff ordered - ordered.roll(-1, dims1) distances diff.norm(dim-1) # (B, N) return distances.sum(dim-1)REINFORCE 损失需要 log 概率。上面的 forward 只返回了 pointers需要额外记录 log_probs。修改 forward 增加返回值# 在 forward 循环内选择 idx 后追加 log_prob torch.log(probs.gather(1, idx.unsqueeze(1)) 1e-9) log_probs.append(log_prob.squeeze(1)) # 循环结束后返回 return torch.stack(pointers, dim1), torch.stack(log_probs, dim1).sum(dim1)训练循环def train_step(model, optimizer, x): model.train() # 采样解码获得一条路径及其 log 概率 pointers, log_prob model(x, decode_typesampling) tour_length compute_tour_length(x, pointers) # 贪心解码作为基线不计算梯度 with torch.no_grad(): greedy_pointers, _ model(x, decode_typegreedy) baseline compute_tour_length(x, greedy_pointers) # REINFORCE 损失最大化 (baseline - tour_length) * log_prob advantage (baseline - tour_length).detach() loss (advantage * log_prob).mean() # 等价于最小化负的期望奖励 optimizer.zero_grad() loss.backward() # 梯度裁剪防止 LSTM 梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() return tour_length.mean().item(), baseline.mean().item()参数说明advantage用baseline - tour_length因为路径越短奖励越高取负号后作为损失方向clip_grad_norm_的max_norm1.0是经验值TSP 训练中梯度范数经常冲到 10 以上不裁剪容易发散。3.2 训练超参与收敛判断完整训练脚本def train(num_epochs100, batch_size256, num_cities20, lr1e-3): device torch.device(cuda if torch.cuda.is_available() else cpu) model PointerNetwork(hidden_dim128).to(device) optimizer torch.optim.Adam(model.parameters(), lrlr) for epoch in range(num_epochs): x generate_tsp_data(batch_size, num_cities, device) sample_len, greedy_len train_step(model, optimizer, x) if epoch % 10 0: print(fEpoch {epoch:3d} | Sample: {sample_len:.4f} | Greedy: {greedy_len:.4f}) return model超参建议表参数推荐值调整方向hidden_dim128N≤20/ 256N20太小欠拟合太大过拟合且慢learning_rate1e-3不收敛降到 1e-4batch_size256显存不够降到 64但梯度噪声变大num_epochs100~200看 greedy 长度是否平稳clip max_norm1.0梯度爆炸时降到 0.5收敛判断看贪心解码的路径长度前 20 个 epoch 快速下降之后缓慢改善。如果 50 个 epoch 后还在震荡检查学习率是否过大或基线是否失效。提示训练时采样解码的路径长度通常比贪心差 5%~15%这是正常的探索代价。如果两者差距超过 30%说明策略方差太大可以增大 batch_size 或改用 rollout baseline。4. 避坑与排查指针网络训练 TSP 最常见的 5 个翻车现场4.1 损失不降反升路径长度爆炸现象训练几个 epoch 后采样路径长度从 4.0 涨到 20 以上贪心解码也同步恶化。原因REINFORCE 的梯度方差过大加上 LSTM 梯度爆炸参数更新方向完全随机。常见触发条件是学习率设成 1e-2 或没有梯度裁剪。解决学习率降到 1e-3 或 1e-4加clip_grad_norm_(max_norm1.0)。如果还不行把基线从贪心解码改成指数移动平均EMA基线平滑效果更好。4.2 模型学会「摆烂」所有路径都指向同一个城市现象指针序列出现大量重复索引mask 机制似乎失效。原因mask 的scatter操作写错了维度或者masked_fill的值不够小比如用了 -1e4 但 logits 量级到了 1e5。另一个可能是解码器初始输入全零导致第一步注意力均匀分布后续陷入局部循环。解决检查mask.scatter(1, idx.unsqueeze(1), 1)中 idx 的形状必须是(B, 1)masked_fill用-1e9而不是-1e4解码器初始输入改用可学习的参数向量不要用全零。4.3 训练集表现好换一组随机城市就崩现象在固定随机种子的 TSP 实例上路径长度 3.8换一组种子变成 6.5。原因模型过拟合到了特定坐标分布。指针网络本身有泛化能力但训练时如果 batch 内城市分布太集中比如都挤在角落编码器学到的特征没有覆盖全空间。解决确保torch.rand生成的坐标覆盖[0,1]×[0,1]全空间每轮重新生成数据不要固定一个 batch 反复训练如果城市数可变训练时混合不同 N 的实例如 15、20、25 交替。4.4 GPU 显存溢出batch_size 降到 1 才能跑现象N50 时 batch_size256 直接 OOM降到 64 还是不够。原因指针网络的注意力分数矩阵是(B, N, H)解码 N 步后中间变量累积。加上 LSTM 的隐藏状态显存占用是 $O(B \cdot N \cdot H)$。解决用梯度累积模拟大 batchbatch_size32累积 8 次梯度再更新等效 batch_size256。或者把编码器换成 Transformer注意力计算可以分块。4.5 贪心解码结果比采样还差现象训练日志里 greedy 长度始终高于 sample 长度。原因模型还没收敛时贪心解码容易陷入局部最优而采样有随机性反而能跳出。这不是 bug是训练早期的正常现象。解决继续训练通常 30 个 epoch 后贪心会反超。如果 100 个 epoch 后仍然如此说明模型容量不够增大hidden_dim或加一层编码器 LSTM。5. 规模化验证与推理加速从 N20 到 N100 的实用技巧训练完 N20 的模型直接拿去做 N50 或 N100 的推理路径长度会明显劣化。指针网络对城市规模有一定泛化能力但需要配合几个技巧。技巧一课程学习Curriculum Learning。不要一上来就训 N50先从 N10 开始每 20 个 epoch 增加 5 个城市直到目标规模。这样编码器逐步适应更长的序列最终 N100 的路径长度比直接训练低 8%~12%。技巧二推理时用 beam search 替代贪心。贪心每步只选概率最大的城市beam search 保留 top-k 条候选路径。k3 时推理时间增加不到 2 倍但路径长度平均改善 3%~5%。实现上维护 k 个解码状态每步扩展后按累积 log 概率排序取前 k。技巧三坐标归一化与尺度不变性。训练时坐标在[0,1]推理时如果输入坐标范围是[0,100]路径长度会放大 100 倍但模型输出的指针序列不变。所以推理前务必把坐标归一化到[0,1]否则注意力分数的数值范围偏移会导致选择错误。验证方法用随机生成的 1000 个 TSP 实例分别跑贪心解码和 beam search统计平均路径长度和标准差。如果标准差超过均值的 15%说明模型在某些分布上不稳定需要增加训练数据多样性。def evaluate(model, num_cities50, num_instances1000, beam_width3): 在随机实例上评估模型 device next(model.parameters()).device model.eval() total_length 0.0 with torch.no_grad(): for _ in range(num_instances // 100): x generate_tsp_data(100, num_cities, device) pointers, _ model(x, decode_typegreedy) lengths compute_tour_length(x, pointers) total_length lengths.sum().item() return total_length / num_instances我自己的习惯是每次改完模型结构先跑 N20 的 100 个 epoch 看收敛曲线确认没有震荡后再上课程学习训 N100。血泪经验是不要跳过小规模验证直接怼大规模否则调参调到最后都不知道是模型问题还是数据问题。希望帮到你。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进