Megatron-LM中GEMM的推导

问题定义:

寻找一种适当的切分 的方式与数据流程,可以得出等价的 以及
 

前向传播的讨论

  • Feature1: 按列拆分
    • 权重矩阵按列拆分时,权重矩阵的列方向数据是完整的,要在一个block上计算完整结果矩阵的数据,则要求 矩阵在行方向上是”完整“的。因此对 矩阵同样做列拆分是无意义的。这一点从矩阵维度的代数表达上也成立,即我们要求 计算时,列维度必须是embed_size。
    • 若将 按行拆分,即, 则将 分布在一块设备上时,只能计算出1/4的数据。
    • 因此,必须将 通信到所有的设备上,使用完整的 与按列拆分的 进行矩阵乘法。即:
  • Feature2: 按行拆分
    • 同样地,对 按行拆分或全复制(不改变列维度),不符合矩阵乘法的要求,不具备意义。
    • 按列拆分时,即 , 可得
      • 可见,行拆分的权重特点是,输入可拆分,输出再聚合计算;而列拆分权重则是输入全复制,输出无计算
  • 因此,矩阵乘法分片的两种基本形式是
      1. 权重按列分片,输入广播,输出同步拼接
      1. 权重按行分片,输入按列分片,输出同步相加
  • Feature3: 按列拆分 后,输出内容的设备分布为按列分布,利用这个特性,也可以进行第二次乘法运算(也可以穿插若干有row locality的运算,如softmax/layer norm等)。简单过程如下
    • 重新定义
    • 第一次运用列拆分分布式乘法后可得,
    • 第二次运用行拆分分布式乘法后可得,

反向传播

  • 按列拆分时,
  • 同理对 按行拆分
  • 二者梯度在分布设备上分别计算即可,无需聚合
 
 
Prev
GPU分布式训练缩写眼保健操:PP、DP、DDP、MP、ZeRO、FSDP
Next
Paged Attention&Orca Paper笔记与解读
Loading...