已合并
修改激活函数重计算描述 #360
AtomGit-Bot创建于 2024年5月27日
修改激活函数重计算描述 #360
已合并
从refs/pull/360/head合入到master
共 3 个文件变更+28-6
| @@ -2,16 +2,39 @@ | |||
| 2 | 2 | ||
| 3 | ## 问题分析 | 3 | ## 问题分析 |
| 4 | 4 | ||
| 5 | -在大模型训练过程中,如gelu激活函数会产生大量的数据,但本身计算量很小。激活函数产生的数据会占据大量内存。 | 5 | +现有的大模型训练框架中,重计算和反向计算是绑定在一起调度的,这严重限制了重计算的灵活性。在某些场景下,会限制重计算在模型性能上的优化。 |
| 6 | + | ||
| 7 | +比如在模型中存在某个流程: | ||
| 8 | + | ||
| 9 | +前向:gelu激活函数模块->后续模块A。 | ||
| 10 | + | ||
| 11 | +反向:后续模块A的反向(需要gelu输出的激活值)->gelu反向(与重计算绑定) | ||
| 12 | + | ||
| 13 | +gelu激活函数会产生大量的数据,但本身计算量很小。此时进行激活函数的重计算可以在性能劣化极少的代价下,减少内存占用。 | ||
| 14 | +但在现有重计算框架下,如果对gelu激活函数模块做重计算,并不能节省gelu函数的输出。这是因为在反向时,模块A所需要的gelu输出的激活值,会早于gelu激活函数模块的重计算流程,所以前向必须保留激活函数的输出,导致激活函数的输出并不能节省下来。 | ||
| 15 | + | ||
| 6 | 16 | ||
| 7 | ## 解决方案 | 17 | ## 解决方案 |
| 8 | 18 | ||
| 9 | -此时进行激活函数的重计算可以在性能劣化极少的代价下,减少内存占用。 | 19 | +本特性重新实现了一套重计算框架,可以将重计算灵活地插入到反向计算之前的任意位置。 |
| 10 | -尤其实在诸如反向传播与梯度计算等操作时,通过仅存储必要的中间结果来节省资源。 | 20 | + |
| 21 | +反向(新框架): | ||
| 22 | + | ||
| 23 | +gelu函数重计算->后续模块A的反向 | ||
| 24 | + | ||
| 25 | +此时,gelu函数的输出已经早于模块A的反向,在前向时就无须保留gelu函数的输出值。 | ||
| 11 | 26 | ||
| 12 | ## 解决思路 | 27 | ## 解决思路 |
| 13 | 28 | ||
| 14 | -设计一种传入激活函数进行重计算的机制,在合适的时机,丢弃重计算模块输出的物理存储,保留逻辑视图。在反向时,利用传入的激活函数重新进行计算,得到结果。 | 29 | +设计一种传入模块函数进行重计算的机制,在合适的时机,丢弃重计算模块输出的物理存储,保留逻辑视图。在反向时,在合适的时机,利用register_hook插入重计算流程。利用传入的函数重新进行计算,得到结果。 |
| 30 | + | ||
| 31 | + | ||
| 32 | + | ||
| 33 | +比如gelu在mlp中的位置如图所示。反向计算需要前向产生的a,b,c, d。其中b, c的shape为(batch, seq , 4hidden_szie),gelu为激活函数,其计算较少,故可将tensor c释放掉,反向在4h->h反向前重新计算。 | ||
| 34 | + | ||
| 35 | + | ||
| 36 | + | ||
| 37 | +在前向4h->h计算完毕后,将c释放,保留逻辑视图。在4h->h grad前,需要将c计算回来。这里使用给d打tensor_hook的方式来进行重计算的插入。 | ||
| 15 | 38 | ||
| 16 | ## 使用场景 | 39 | ## 使用场景 |
| 17 | 40 | ||
| @@ -31,8 +54,7 @@ | |||
| 31 | 54 | ||
| 32 | ## 扩展使用 | 55 | ## 扩展使用 |
| 33 | 56 | ||
| 34 | -在设计传入激活函数进行全重计算时,引入的 CheckpointWithoutOutput 类具有通用进行函数重计算的机制。 | 57 | +本特性引入的 CheckpointWithoutOutput 类可以自定义对任何模块进行重计算,并且在合适的时机进行重计算恢复。 |
| 35 | -不仅可以对激活函数进行重计算,也可以对任何自定义的函数进行重计算。 | ||
| 36 | 58 | ||
| 37 | 此处提供一个示例,可以灵活使用 CheckpointWithoutOutput 来对自定义的函数进行重计算: | 59 | 此处提供一个示例,可以灵活使用 CheckpointWithoutOutput 来对自定义的函数进行重计算: |
| 38 | 60 | ||