解题思路
核心思路
本题是在显存预算约束下,选择若干中间层打 Checkpoint,使重计算总代价最小。相邻两个 Checkpoint(含隐式起点 0 与终点 N)之间的代价为
Cost(c,d)=k=c+1∑d−1forward_time[k−1]⋅(d−k)
可先用前缀和 S,由递推 Cost(c,d)=Cost(c,d−1)+(S[d−1]−S[c])、Cost(c,c+1)=0 在 O(N2) 内预处理全部区间代价。
题目描述
在大模型训练中,前向传播会留下大量中间激活,反向传播时还要用到它们。当 GPU 显存吃紧时,激活重计算(Gradient Checkpointing)是一种常见的「用时间换空间」做法:只在若干「检查点」层保存激活,其余层在反向阶段再算一遍。本题要求:在给定显存预算下,求出最优的 Checkpoint 放置方案。
模型共有 N 层(编号为 1 到 N),按前向顺序依次执行。给定两个长度为 N 的正整数数组:
forward_time[i](0≤i<N):第 i+1 层前向传播的耗时;
memory[i](0≤i<N):在第 i+1 层保存激活(打 Checkpoint)所需的额外显存。
固定边界:
- 第 0 层(输入层)的激活常驻显存,不计入额外开销,它是天然的起始 Checkpoint;
- 第 N 层是终点:不必在第 N 层再存额外激活,也不占用 Checkpoint 显存。
操作:你可以在某些中间层 c(1≤c≤N−1)打 Checkpoint。在第 c 层打 Checkpoint 会占用 memory[c-1] 的显存(第 c 层对应数组下标 c−1)。所有 Checkpoint 的显存占用之和不得超过 MaxMem。
重计算代价:
设两个相邻 Checkpoint 位置为 c 与 d(c<d,中间没有别的 Checkpoint)。反向传播经过该区间时,需要从第 c 层的激活出发,重新前向计算以恢复中间层激活。由于反向是逐层推进的,层 d−1 只需重算一次(从 c 前向到 d−1),层 d−2 需要重算两次,以此类推。因此该区间的重计算代价为:
Cost(c,d)=k=c+1∑d−1forward_time[k−1]×(d−k)
其中 forward_time[k−1] 是第 k 层的前向耗时,(d−k) 是该层被重复计算的次数。若 d=c+1(相邻 Checkpoint),代价为 0。
目标:在显存预算 MaxMem 内选择 Checkpoint 放置方案,使整网总重计算代价最小。总代价等于所有相邻 Checkpoint 对之间的代价之和(隐含起点 0 与终点 N)。
输入描述
输入格式如下:
N MaxMem
f_1 f_2 ... f_N
m_1 m_2 ... m_N
- 第 1 行:两个非负整数 N 与 MaxMem,分别表示模型层数与额外显存预算上限。
- 第 2 行:N 个空格分隔的正整数,表示每层前向传播耗时。
- 第 3 行:N 个空格分隔的正整数,表示在每层保存 Checkpoint 所需的额外显存。
保证:1≤N≤100,0≤MaxMem≤1000,1≤fi,mi≤1000。
输出描述
输出一行一个整数:最小总重计算代价。
样例1
输入
5 3
5 10 5 10 20
1 2 1 3 1
输出
15
说明
可选的 Checkpoint 位置:层 1,2,3,4。始终存在隐式起点 0 与终点 5。
| 方案 |
Checkpoint 位置 |
显存消耗 |
区间划分 |
各区间代价 |
总代价 |
| A |
无 |
0 |
(0,5) |
5⋅4+10⋅3+5⋅2+10⋅1=70 |
70 |
| B |
{2} |
2 |
(0,2),(2,5) |
5,5⋅2+10⋅1=20 |
25 |
| C |
{1,3} |
(0,1),(1,3),(3,5) |
0,10,10 |
20 |
| D |
{2,3} |
3 |
(0,2),(2,3),(3,5) |
5,0,10 |
15 |
方案 D 显存消耗 3≤3,总代价 15 最小。
Cost(0,5) 细算:
- k=1:forward_time[0]×(5−1)=5×4=20
- k=2:forward_time[1]×(5−2)=10×3=30
- k=3:forward_time[2]×(5−3)=5×2=10
- k=4:forward_time[3]×(5−4)=10×1=10
合计 70。注意第 5 层(k=5)不在求和范围内(上界为 d−1=4)。
补充说明
提示
可建模为带容量约束的区间最短路 / 分段 DP。
先计算 cost[c][d](0≤c<d≤N),利用前缀和可在 O(N2) 内完成:
- 先算 S[i]=∑k=1iforward_time[k−1](前缀和)
- Cost(c,d)=∑k=c+1d−1forward_time[k−1]⋅(d−k),可用递推在 O(1) 得到每对。
再构建转移:对每个位置 i 与已用显存 w,枚举上一个 Checkpoint 位置 j(0≤j<i):
dp[i][w]=jmin(dp[j][w−memory[i−1]]+cost[j][i])
其中 dp[0][0]=0(起点不耗显存),第 N 层作为终点不消耗额外显存。