CUDA算子优化(2):RMSNorm
1. 简介
RMSNorm和Batch Normalization(BN)、Layer Normalization(LN)等方法一样,都属于一种归一化方法,是提升训练稳定性、加速收敛的重要技巧之一。这里,我们不讨论这些方法的优劣性,我们只关注RMSNorm是如何计算的,以及CUDA程序如何实现和优化。
RMSNorm仅使用均方根归一化,即计算数据的平方的平均值,再开平方根,用来衡量数据的整体大小。核心公式为输入数据除以他们的均方根,然后乘以一个可学习的缩放参数 $\gamma$ 。具体公式可表示为:
$$
\bar{x}_i = \frac{x_i}{\mathrm{RMS}(x)} \cdot \gamma_i
$$
其中,$\mathrm{RMS}(x) = \sqrt{\frac{1}{C} \sum_{i=1}^{C} x_i^2+ \epsilon}$
2. 算子CPU实现
假设我们要处理的序列维度是B=8,T=1024,C=4096,其中B是batch_size,T是序列的长度,C是通道数;要求均方根,是在C这个维度进行求解。假设h_data为输入,维度为B*T*C;h_weights为缩放因子,维度为C;h_out为输出,维度为B*T*C。我们把计算过程分为三个阶段:
- 计算每一行的平方和
- 计算RMS(x)
- 乘上缩放因子,计算最终结果
|
3. CUDA版本V0一个线程处理一行
假设一个线程处理一行,那么总共需要B*T个线程,以block_size=64为例,grid_size为(B*T + block_size - 1) / block_size。代码如下:
|
4. CUDA版本V1一个Warp处理一行
很明显,可以看出,使用一个线程处理一行访存是不合并的,所以我们可以使用一个Warp处理一行,减少访存不合并的现象;总共需要B*T*warp_size个线程。我们仍然将其分为三个阶段。
- 计算每一行的平方和。相比于使用共享内存,我们可以直接使用__shfl_down_sync函数,这个函数直接在寄存器层面完成Warp内线程间数据交换,完全规避了共享内存的硬件延迟、同步开销和Bank Conflict的风险,速度更快。这里介绍一下__shfl_down_sync和__shfl_sync函数
__shfl_down_sync(mask, v, d, w),标号为t的参与线程返回标号为t+d的线程中变量v的值,标号满足t+d>=w的线程返回原来的v。例如,当w=8,d=2时,该函数将第2-7号线程中变量v的值传递给第0-5号线程,而第6-7号线程返回它们原来的v。形象地说,这是一种将数据向下平移的操作。
__shfl_sync(mask, v, srcLane, w),参与线程返回标号为 srcLane 的线程中变量 v 的值。这是一种广播式数据交换:将一个线程中的数据广播到所有(包括自己)线程。
Warp间数据规约的代码如下:
|
- 计算RMS(x)
- 乘上缩放因子,计算最终结果
全部代码如下:
|
5. CUDA版本V2一个Block处理一行
当数据量比较大的时候,一个Warp处理一行,显然是不够的,这个时候,我们就可以使用一个Block处理一行,总共需要B*T*Block_size个线程,我们仍然将其过程分为三个阶段。
- 计算均方根。这里我们就不能使用__shfl_down_sync等洗牌函数了,因为这个函数是Warp间的规约,不适用于Block,所以,我们只能引入共享内存。使用共享内存的方法已经在第一节CUDA算子优化(1):Reduce中讲过,请参考第一节。
- 计算RMS(x)
- 乘上缩放因子,计算最终结果。
|