PyTorch中结合滑动窗口与Hirschberg算法的0/1背包问题求解的内存高效激活检查点
Memory-Efficient Activation Checkpointing with Sliding Window and Hirschberg's Algorithm for 0/1 Knapsack Solving in PyTorch
浏览论文内容
中文总结 AI 辅助
该研究针对PyTorch激活检查点内存效率问题,提出结合滑动窗口与Hirschberg算法的0/1背包求解方法,可处理规模提升20倍,速度提升25%-28%,已合并入PyTorch 2.10版本。
中文摘要 AI 辅助
激活检查点通过选择存储哪些中间张量、重计算哪些中间张量,在给定内存预算下最小化神经网络的运行时间。PyTorch将该问题建模为0/1背包问题,其中联合前向-后向计算图中的操作作为物品,其内存开销对应背包重量,运行时间节省对应背包价值。默认求解器dp_knapsack会分配形状为(n+1)×(W+1)的完整动态规划(DP)表,其中n为操作数量,W为量化后的内存预算,该方法资源消耗大,在64GB RAM的机器上处理n=100个物品时会崩溃。本文提出dp_knapsack_sliding_hirschberg,结合滑动窗口技巧与Hirschberg算法,在保留精确最优解的同时将峰值内存从O(nW)降至O(W)。实验显示,该方法可成功处理n=2000个物品的背包问题,而dp_knapsack在n=100时即失败,可处理的问题规模提升20倍;此外,基准测试显示其运行时间比dp_knapsack快25%-28%,该实现已合并入PyTorch并在2.10版本发布。
英文摘要
Activation checkpointing minimizes the runtime of neural networks under a given memory budget, by selecting which intermediate tensors to store and which to recompute. PyTorch solves this as a 0/1 knapsack problem, where operations from a joint forward-backward computation graph are items with a memory cost (weight) and a runtime saving (value). The default solver, dp_knapsack, allocates a full dynamic programming (DP) table of shape $(n+1) \times (W+1)$, where $n$ is the number of operations and $W$ is the quantized memory budget. This method is resource-hungry and crashes at $n = 100$ items on a machine with 64 GB RAM. In this paper, we introduce dp_knapsack_sliding_hirschberg, which combines the sliding window trick and Hirschberg's algorithm to reduce peak memory from $O(nW)$ to $O(W)$ while preserving the exact optimal solution. Our experiments show successful knapsack execution at $n = 2000$, where dp_knapsack fails at $n = 100$, a 20$\times$ increase in computable problem size. In addition, our benchmarks show a consistent 25-28\% runtime speedup over dp_knapsack. The implementation is merged into PyTorch and released in version 2.10.
发表机构
- Cohere Labs(Cohere实验室)
机构由 AI 辅助整理,请以论文原文为准。