ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

【CS336】lecture5 Roofline model|线程束分化|低精度|算子融合|重计算|合并访存|分块

【CS336】lecture5 Roofline model|线程束分化|低精度|算子融合|重计算|合并访存|分块 概览这部分的主题是如何在GPU上优化AI任务。开始这个图就很有代表性表示了随着矩阵规模增长矩阵乘法的吞吐如何增长可以发现有大趋势也有奇怪的细节后面会一一解释。Roofline model一个经典的分析是roofline model屋顶模型随着计算强度增长GPU吞吐呈现一个屋顶形态的曲线。其中横坐标计算强度的定义是计算量/搬运量。这是我们之前提到的内存带宽的发展落后于计算核心的发展内存带宽普遍比计算核心的吞吐低至少一个数量级因此想要让GPU吞吐打满必须给计算单元足够的数据。如图GPU的吞吐在计算强度10以上才能达到峰值也就是搬运1单位数据需要计算10次以上才能让计算单元一直处于满载状态。如果搬运1次就计算1次的话大部分时间计算单元都是闲置的时间都花在等待数据搬运。同时访问的内存在不同的位置打满吞吐的转折点也不一样访问的内存越快打满吞吐需要的计算强度越低。比如寄存器是最快的他的转折点计算强度甚至不到0.1也就是内存带宽是计算单元吞吐的10倍以上。共享内存SMEM适中转折点在1-10这个区间内全局内存GMEM是最慢的需要10-100的计算强度才能打满优化思路接下来将探讨主要的几种GPU优化思路线程束分化低精度运算算子融合重计算内存对齐分块线程束分化先回忆之前的执行模型很多线程调度器每次会把32个线程绑一块成为一个线程束warp一块执行那么最好情况下一个warp内32个线程做的操作应该完全一样这样32级并行。但如果我们的代码里有控制流比如一个if-else可能导致一部分线程走了if一部分走了else那么这两部分内部是可以并行的但两部分之间是串行的这被称为线程束分化Control divergence拖慢执行速度因此最好在代码里避免这种可能导致分化的if-else结构低精度低精度的好处是搬运量降低了计算量不变的话计算强度能提高。比如这里把fp32换成fp16relu算子的计算强度从1/8提升到1/4图中是取了个倒数根据前面的roofline model计算强度提升是好事。此外低精度还能让计算更快比如tensor core做矩阵乘只在累加这个需要精度的地方用fp32AB矩阵输入都用fp16这样单位面积的芯片能容纳更多计算单元计算吞吐提高。也就是之前roofline model的屋顶计算吞吐上界可以提高。FP8低精度最前沿的进展是8bit浮点数。在位数位和精度位上有多种实现主流是E4M34位指数3位位数第0位是符号位。更新的处于理论阶段的是MXFP8。FP8在使用时需要缩放比如我们把[-100,100]映射到[-1,1]的范围内需要除以一个缩放因子100。对于传统FP8全局维护一个缩放因子即可。对于MXFP8对数据做了分块每块维护一个缩放因子如图右侧就是MXFP8每一种颜色表示一个块一个颜色的缩放因子是下方相同颜色的因子。这里缩放因子也采用FP8只不过是E8M0也就是8位全部用来存指数位因为缩放因子类似科学计数法里的E10的这个10不需要尾数这样的好处是不同数据可以有不同的缩放尺度低精度量化更加灵活。问题是矩阵转置会变复杂不仅需要转置原始数据还要转置缩放因子矩阵。按这个思路还可以有FP4如图指数位2位位数1位当然这样表示的不同数值就只有16个了。算子融合首先是一个比喻计算单元像个工厂内存带宽是搬运原材料和产品的公路工厂产能拓展的很快但公路运力拓展的很慢。这其实就是我们前面说过的计算单元吞吐增长比内存带宽增长快至少一个数量级。当要对一个产品多多重加工时朴素的做法是每次加工完一道工序搬回内存然后再发往工程做下一道工序如图左侧三道工序的话要搬运6次。既然瓶颈在于搬运那么一个优化思路是让三道工序再工厂一次性做完再搬回来这样只需要2次搬运这就是算子融合。具体来说就是把多个算子的工作合并到一个算子里比如这里要计算sin⁡2xcos⁡2x\sin^2x\cos^2xsin2xcos2x朴素的方法需要5次算子调用分别是sin⁡x,cos⁡x\sin x,\cos xsinx,cosx两次平方一次加法。每次调用都要把数据从全局内存搬运到计算单元的专用内存SMEM,寄存器实际上完全可以融合到一个算子里这样只需一次调用数据搬运次数大大减少。大部分算子尤其是element-wise逐元素操作算子都可以这样优化在编程中调用torch.complie就会自动做算子融合重计算在反向传播计算梯度时需要每一层的激活值一般的做法是把前向传播每一层的结果缓存下来但这导致了大量的内存读写比如这个三层sigmoid前向传播时要保存三层的激活值写内存3次反向传播时需要激活值又要读3次。而sigmoid的计算量是很小的于是计算强度很低。所以现在更多的做法是反向传播时重新计算前向传播的激活值比如这个图中的流程前向传播除了最终的输出不额外保存。反向传播时再跑一次前向传播计算出每一层的激活值这样内存读写次数大幅降低计算强度提升。尽管计算量增加了但是注意现代GPU的计算吞吐远大于内存带宽这点计算的耗时远比内存搬运的耗时要短。当然实际的策略也不是全部重新计算而是在计算和访存间做一个取舍比如每隔d层保存一次激活值也就是激活值分块保存每块保存块内开头层的激活值反向传播时从这一块的开头开始重新计算。这就是所谓的检查点机制合并访存根据内存的硬件特性每次访问内存都是访问一块而不是只访问某一个位置例如这个图里16个位置划分成4块每次读都是读其中一块。也就是我只想读1的话也会把0123都读出来。在实际中国模当然比这个要大但这个图演示的规律还是成立的。一般内存4GB单次读128B。基于这个特性我们在编程时应该尽量让一个warp内的线程访问连续的内存比如4个线程一组的话我们让这4个线程访问一块的连续4个内存位置就只需要一次内存事务。如4个线程分别需要04812的话则需要四次内存事务耗时是前面的4倍。这是经典例子读一个二维矩阵如果我们让每个线程负责一行那么就不是内存合并的因为一般的矩阵在内存中保存的方法是行主序那么两个线程分别管两行实际上他们访存的位置跨越了一行的长度。Tiling这可以说是最重要的优化。考虑一个矩阵乘法关注左上角的子矩阵分配给四个线程执行可以发现访问的很多位置是重复的比如M(0,0)在线程(0,0)(0,1)中都访问了我们又知道全局内存慢共享内存快于是一个自然的想法就是把数据分块搬运到共享内存里接下来计算时频繁读取都直接读共享内存。一次完整的矩阵乘原来的访存方式下输入矩阵的每个元素都要被读N次也就是N次全局内存读。现在Tiling后假设分成T组那么每个数据会有T次全局内存读也就是搬运到T组分块的共享内存里接下来在块内都是共享内存读了。也就是可以把最耗时的全局内存读优化N/T倍。最终一个经典的Tiling矩阵乘法的过程是把结果数组分块得到C tile也就是橙色的部分每一个结果矩阵块需要输入矩阵AB的两块也就是浅紫色的两块搬运到共享内存中在AB tile上做一个划窗每个划窗做矩阵乘法累加到C tile上每次矩阵乘法是一个O(d3)O(d^3)O(d3)的循环d是块的大小在这图中就是4。图中绿色部分就是表示这部分循环对于C tile里的一个绿色位置需要AB tile里的一行和一列。在tiling时分块大小也有讲究需要考虑以下几点合并访存有点复杂后面细说共享内存大小这是因为我们要把AB tile都加载进共享内存如果tile太大共享内存放不下块长能否整除矩阵大小如果不整除边界块的任务量是比中间整块要小的造成了资源的浪费分块大小实际上还会影响合并访存。注意这图左侧的tile大小正好可以整除一行的长度3个tile是一行同时也等于合并访存一次的长度也就是取第一个列tile块时在四行每行做一次合并访存即可。右侧的tile大小则不好3个tile还要再多一点才是一行的长度所以3tile一起还不够拼出第一行的蓝色让蓝色拐歪在第二行还剩下一点。这导致我们想取第一个列tile时也就是取(M,N_tile)大小的列块时四种颜色的列块都插在两个合并访存块中间四行每行要做两次访存事务才能取出第一个列块。举例karpathy的优化案例他发先把矩阵大小从50257对齐到最近的64的倍数50304就带来了25%的性能优化。在大部分时候2的幂次总是好的因为合并访存bank conflict等优化都是2的幂次对齐的。50257是他原本的词表真实大小他填充了一些无意义的字符或者空字符拓展到50304这确实导致了一些冗余计算但可能优化了访存正如前面所说的GPU上计算量往往不重要很少是计算bound重要的是吞吐也就是按roofline模型的思路提高计算强度打满吞吐。回到开始现在我们能解释这张图了。矩阵乘法的计算强度关于矩阵规模N计算量是N3N^3N3搬运量只有N2N^2N2因此计算强度是O(N)O(N)O(N)的所以随着矩阵规模的提升计算强度越来越高根据roofline model吞吐也在提升直到计算强度够大后达到屋顶。这是这张图的大趋势另一个趋势是在相同矩阵规模下选择不同的tiling分块长度也对吞吐有很大的影响这是我们前面讨论过的分块大小会影响合并访存使用的共享内存和寄存器进而影响SM占用率。所以tile不是越大越好也不是越小越好这个图反应最优的是16其次是8再往下是2最后是32仔细观察的话还会发现这个图不是单增的同一种颜色也就是相同的tiling中当矩阵规模增长时吞吐有时会出现反常的下降。尤其是1792-1793这里右侧的计算分析告诉我们1792的时候会分出98个C tile每个C tile对应一个线程块交给一个SM计算实验用的A100有108个SM足够装下这些tile并且一轮并行跑完。1793时为了上取整额外启动了很多tiletile总数来到120超过A100一次的最大并行个数于是有12个tile被留到第二轮运行由于同步约束第二轮12个SM工作时剩下96个SM也只能干等着造成了吞吐的极大浪费。
返回列表