应该是一个比较常见的手段了,之前在Megatron MoE的paper中他们有系统的讲。今天在看recompute相关的东西的时候突然想到了这个,就来整理一下。

逻辑是这样的,在MoE fc2之后,我们有一步根据probs加权的操作。
考虑到,对于一个matmul(fc2)来说,他的backward只依赖他的输入,不依赖他的输出。
也就是说如果没有最后那个加权,FC2的输出是不需要save的。
这里就结合到矩阵乘法的特性,g * (h @ W2) = (g * h) @ W2。只要满足g是在行这个维度的放缩即可。
直观理解的话:
- 考虑输出的元素o(i, j),他是 for k,a(i, k) * (b(k, j)累加。
-
如果a(i)行有一个相同因子的放缩,就可以提出来。即 scale(i) * sigma(a(i, k) * b(k, j))
-
因为每一行都一样,自然就可以提到output外面,变成scale(i) * o(i, j)了。
同理列也是一样的。只要放缩的维度不是k就行。
这样的话,在MoE的时候,可以把probs也做通信,相比于hidden通信量很少。然后在fc2之前先做放缩。再做fc2。
这样就可以省掉fc2的output了。不过fc2因为是down proj,所以hidden比fc1小,整的比例不多。但是是无痛的。
同样,对于recompute来说也是一样的,fc2不需要做重算,因为backward不依赖这个fc2的output。这个应该可以被torch AC中的early return给处理掉。
- 当然也依赖用户的写法了。最好是tensor准备好了就立刻save for backward。
文章评论