More than code

More Than Code
The efficiency of your iteration of reading, practicing and thinking decides your understanding of the world.
  1. 首页
  2. 未分类
  3. 正文

MoE fc2 显存优化

2026年7月12日 64点热度 0人点赞 0条评论

应该是一个比较常见的手段了,之前在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。
标签: 暂无
最后更新:2026年7月12日

sheep

think again

点赞
< 上一篇

文章评论

取消回复

COPYRIGHT © 2021 heavensheep.xyz. ALL RIGHTS RESERVED.

THEME KRATOS MADE BY VTROIS