[论文解读] Differentiation of Blackbox Combinatorial Solvers
该论文提出了一种针对具有线性目标函数的黑箱组合优化求解器的可微反向传播方法,实现了将精确求解器(如 Gurobi、Blossom V 和 Dijkstra 算法)作为可微组件嵌入神经网络的端到端训练。该方法通过离散解的连续插值计算梯度,仅需额外一次求解器调用即可实现精确的反向传播,从而支持混合模型的训练,能够以当前最先进求解器的性能解决复杂问题,如旅行商问题(TSP)、最小费用匹配和最短路径问题。
Achieving fusion of deep learning with combinatorial algorithms promises transformative changes to artificial intelligence. One possible approach is to introduce combinatorial building blocks into neural networks. Such end-to-end architectures have the potential to tackle combinatorial problems on raw input data such as ensuring global consistency in multi-object tracking or route planning on maps in robotics. In this work, we present a method that implements an efficient backward pass through blackbox implementations of combinatorial solvers with linear objective functions. We provide both theoretical and experimental backing. In particular, we incorporate the Gurobi MIP solver, Blossom V algorithm, and Dijkstra's algorithm into architectures that extract suitable features from raw inputs for the traveling salesman problem, the min-cost perfect matching problem and the shortest path problem. The code is available at https://github.com/martius-lab/blackbox-backprop.
研究动机与目标
- 实现将精确、未经修改的组合优化求解器作为可微组件嵌入神经网络的端到端训练。
- 克服通过分段常数、不可微的组合优化求解器(具有线性目标)进行反向传播的根本性挑战。
- 在保持高性能求解器(如 Gurobi、Blossom V)最优性和性能的同时,实现基于梯度的优化。
- 提供一种适用于广泛问题的通用方法,包括 TSP、最小费用完美匹配、最短路径和整数规划问题。
提出的方法
- 通过扰动权重下的最优解的凸组合,对离散解进行连续插值。
- 利用组合问题的最小化结构以及最优值函数的次微分,推导出损失相对于输入权重的梯度。
- 采用基于损失增强推理和分段常数插值的新型数学技术,无需修改求解器即可计算梯度。
- 反向传播仅需额外调用一次黑箱求解器,计算成本与前向传播完全一致。
- 该方法适用于任何具有线性目标函数的组合问题,包括 TSP、最小费用匹配和最短路径问题。
- 通过超参数 λ 控制插值,确保梯度信息量充足的同时保持求解器的保真度。
实验结果
研究问题
- RQ1我们能否在不改变其内部逻辑或性能的前提下,对未经修改的黑箱组合优化求解器(具有线性目标)实现反向传播?
- RQ2所提出的方法是否在端到端训练中保持了高性能求解器(如 Gurobi 和 Blossom V)的最优性和运行效率?
- RQ3嵌入精确求解器的混合神经架构能否解决标准深度学习模型难以处理的复杂组合任务(例如地球表面的 TSP)?
- RQ4使用近似求解器(如 OR-Tools)与使用精确求解器相比,该方法的性能如何?
- RQ5超参数 λ 对梯度质量和训练稳定性有何影响?
主要发现
- 该方法实现了将精确求解器(如 Gurobi、Blossom V 和 Dijkstra 算法)嵌入神经网络的端到端训练,在 TSP、最小费用匹配和最短路径任务上达到最先进性能。
- 在 Globe TSP 基准测试中,对于 k=40 个城市,模型平均城市定位误差为 58 ± 7 km,经 Procrustes 变换后最佳结果与真实值高度一致。
- 在 MNIST 最小费用完美匹配任务中,采用全卷积网络和 λ=10 时,k=8 和 k=16 的准确率均达到 100%。
- 在 Warcraft 最短路径数据集上,k=30 时模型测试准确率达到 99.3%,表明其对复杂真实地图结构具有强大的泛化能力。
- 当使用近似求解器 OR-Tools 时,k=10 时模型测试准确率达到 84.4%,接近在输入真实位置时的上限 88.6%,表明该方法对求解器次优性具有鲁棒性。
- 该方法保持了计算效率,反向传播仅需额外一次求解器调用,计算成本与前向传播完全一致。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。