首页 / 资讯中心 / 文章详情

【CS336】lecture6 占用率|bank冲突|Triton|GeLU算子|softmax算子|reduce算子|matmul算子

【CS336】lecture6 占用率|bank冲突|Triton|GeLU算子|softmax算子|reduce算子|matmul算子 ★ FEATURED ARTICLE
回顾这节主要讲的是如何用triton写算子优化性能。首先回顾一下上一节讲的GPU架构。一个GPU上有多个SM多个SM共享全局内存HBM以及一个全局缓存L2。每个SM共享L1缓存和共享内存。每个线程独享寄存器资源但一个SM上的寄存器其实也是公用的一块区域划分出来的。AcceleratorA100H100B200SM 数量108132148Register 大小每个 SM256 KB256 KB256 KBL1 cache shared memory每个 SM192 KB256 KB256 KBL2 cache 大小40 MB50 MB96-126 MBHBM 大小80 GB80 GB192 GBRegister bandwidth~116 TB/s~401 TB/s~447 TB/sL1 cache shared memory bandwidth~19 TB/s~33 TB/s~19 TB/sL2 cache bandwidth~5-8 TB/s~12 TB/s~9 TB/sHBM bandwidth2 TB/s3.35 TB/s8 TB/s对于各级内存的具体大小和带宽如上表。对于线程模型每次启动只能启动一个gridgrid的shape决定了线程块block的排布。每个block也称为CTA(Cooperative Thread Array)保存了一组线程。最底层是单个线程thread。对于逐元素操作看起来让每个线程各干各的就好了不需要block这层抽象。设置block这层的原因是有些算子需要块内协作比如softmax需要求块内最大值求和才能计算每个位置的结果也就是每个线程需要访问其他线程的数据这如果都通过全局内存就太慢了所以设置block每个block内共享SMEM可以把块内共享的数据加载到SMEM中快速访问。Warp occupancy上节课没提到的GPU优化的另一个重点。在一个 thread block 内thread 被分为 warp每个 warp 有 32 个 thread。例如一个有 64 个 thread 的 thread block 包含 2 个 warp。每个 thread 可以使用 0 到 255 个 register。每个thread 使用的寄存器越多一个 SM 能调度的 thread 就越少即 occupancy 越低。如果每个 thread 做了更多工作低 occupancy 不一定是坏事。一个例子是 thread coarsening即每个 thread 处理多个 element。但申请了不必要的过多的寄存器会导致能并发的warp数太少。一个SM并发的warp数多有利于SM调度可以在一个warp等待数据读写时调度另一个warp进行计算类似于CPU里的超线程技术并发线程数大于物理核心数可以提高SM/CPU的利用率。对于warp occupancy的严格定义是当前warp数和SM的最大warp数的比值一般来讲越高越好。当然除了寄存器限制SM利用率的因素还有共享内存一个SM会有多个block每个block都需要共享内存而SM的共享内存有限会限制block个数进而限制了线程数和warp个数。Bank conflict另一个上节课没提到但很重要的东西是bank冲突。bank指的是共享内存实际被划分为32个可并行读写的块正好等于warp大小在设计上理想情况下warp内每个线程都去读写共享内存的不同bank实现完全的并行内存访问。如果32个线程不是访问32个不同bank而是有些线程会访问同一个bank此时在这个多线程都要访问的bank上请求会退化成串行处理这被称为bank conflict。由于warp是同步的所有线程都要等待这个bank的几个线程串行读写完毕访存延迟大幅增加如图最左侧的情况是线程和bank一一对应不存在bank冲突。中间的图大部分一一对应有几个bank出现冲突访存延迟略微增加。最右边的情况32个线程只读2个bank出现了严重的bank冲突访存延迟大幅增加。对此的主流解决方法是swizzling把数据按照一定规则打乱让一个warp内的线程读取的是不同bank的数据Triton这节给一些算子的实际例子用triton实现triton是一个类python的GPU算子开发语言写类py语法然后triton编译器会帮你编译成ptx汇编也就是GPU的专用汇编代码因此可以获得媲美纯CUDA的执行效率同时还有python的开发效率综合来讲是一个比较好的语言。在相当长的时间内triton都是唯一的主流算子开发DSL领域专用语言vllmSGLang甚至OpenAI内部有大量算子是用triton实现的。直到近几年才有mojo,tilelang,cutile等竞品开始挑战triton的地位。triton简化了编程模型编程者不再能直接控制线程而只能控制到block这个维度这虽然损失了一定的灵活性和性能上限但是简化了开发。开发时只用考虑如何把数据分块每一块要做什么计算即可不用考虑每个线程需要做什么把对初学者可能造成困惑的SIMT编程范式替换成了便于理解的SIMD。如上图左侧是CUDA编程模型需要考虑每个线程做什么每个线程具体负责哪个元素右侧是triton只用考虑划分成那几个数据块每个数据块需要做什么相当于在CUDA中只需要考虑block不考虑threadwarp这些。GeLUdeftriton_gelu(x:torch.Tensor):assertx.is_cudaassertx.is_contiguous()ytorch.empty_like(x)num_elementsx.numel()BLOCK_SIZE1024num_blockstriton.cdiv(num_elements,BLOCK_SIZE)kerneltriton_gelu_kernel[(num_blocks,)](x,y,num_elements,BLOCK_SIZEBLOCK_SIZE)output_ptx(triton_gelu,kernel)returnytriton.jitdeftriton_gelu_kernel(x_ptr,y_ptr,num_elements,BLOCK_SIZE:tl.constexpr):pidtl.program_id(axis0)startpid*BLOCK_SIZE offsetsstarttl.arange(0,BLOCK_SIZE)maskoffsetsnum_elements xtl.load(x_ptroffsets,maskmask)# Approx gelu is 0.5 * x * (1 tanh(sqrt(2/pi) * (x 0.044715 * x^3)))# tl.tanh doesnt exist; use tanh(a) (exp(2a) - 1) / (exp(2a) 1)a0.79788456*(x0.044715*x*x*x)exptl.exp(2*a)tanh(exp-1)/(exp1)y0.5*x*(1tanh)tl.store(y_ptroffsets,y,maskmask)从一个例子来看利用triton的算子开发。triton_gelu()是对外提供的python函数接口启动kernel之前的准备工作都在这里做包括确定block大小block大小必须是2的幂次这里直接硬编码1024根据block大小计算分块个数检验数据是否传送到显存上检验数据排布是否连续。triton_gelu_kernel()是triton内核从这个内核可以初步了解一下triton的编程范式。首先确定当前内核的编号pid根据这个编号确定当前block对应的数据区间也就是从下标start开始长度BLOCK_SIZE的一块数据。根据这个区间构造一个下标数组offsets保存这个区间内的下标。为了防止越界再构造一个mask数组和offsets等长每个位置表示offsets里的对应下标是否有效不能超出num_elements后面就是体现triton的block思想的地方了我们直接用offsets和mask数组作为参数调用tl.load()搬运当前block负责的数据块不必像CUDA中需要计算每个线程负责的下标而是直接整块搬运。后面对这取出来的这个数据快当成类似numpy数组来处理利用py的语法糖对整块数据做统一操作也就是GeLU的具体定义最后还是用tl.store()和tl.load()类似也传入offsets和mask。不同的是写入结果地址是y指针开头的地址不是输入时的x指针对于运行triton只是长得像py但他并不是py的解释执行而是和CUDA一样也会编译成完整PTX汇编直接在GPU上运行可执行程序这也是triton性能可以媲美纯CUDA的原因。并且事实上triton代码往往很简单很多优化工作是编译器自动帮我们做的triton的编译器优化的很好随手写的简单triton实际上往往就可以打败大多数手写CUDA了只是和顶尖CUDA优化有差距。softmaxdeftriton_softmax(x:torch.Tensor):ytorch.empty_like(x)M,Nx.shape block_sizetriton.next_power_of_2(N)# Each block contains all columns num_blocksM # Each block is a row triton_softmax_kernel[(M,)](x_ptrx,y_ptry,x_row_stridex.stride(0),y_row_stridey.stride(0),num_colsN,BLOCK_SIZEblock_size,)returny triton.jit deftriton_softmax_kernel(x_ptr,y_ptr,x_row_stride,y_row_stride,num_cols,BLOCK_SIZE:tl.constexpr,):assert num_colsBLOCK_SIZE row_idxtl.program_id(0)col_offsetstl.arange(0,BLOCK_SIZE)x_start_ptrx_ptrrow_idx*x_row_stride x_ptrsx_start_ptrcol_offsets x_rowtl.load(x_ptrs,maskcol_offsetsnum_cols,otherfloat(-inf))x_rowx_row-tl.max(x_row,axis0)numeratortl.exp(x_row)denominatortl.sum(numerator,axis0)y_rownumerator/denominator y_start_ptry_ptrrow_idx*y_row_stride y_ptrsy_start_ptrcol_offsets tl.store(y_ptrs,y_row,maskcol_offsetsnum_cols)GeLU是逐元素算子再来看个复杂一点的逐行规约算子。分块规则是每个block负责一行。列数是变长的所以块大小不能硬编码了应该取不小于列数的最小的2的幂。输入是一个矩阵kernel的内的寻址稍微复杂一点。首先找到所在行的开始地址x_start_ptr。然后构造一行内的下标数组col_offsets也就是一个[0,BLOCK_SIZE)的列表。col_offsets 和x_start_ptr 加起来得到这一行的具体下标数组x_ptrs用这个数组来做tl.load。mask设计为不超过列数的为1别的为0softmax计算时进行幂操作想让无效位置不影响结果应该赋值为负无穷接下来做一下safe softmax的操作找到行内最值给每个元素减去最值。求以e为底数的幂。规约求和作为分母除以幂。都可以用类numpy/torch.tensor的重载快速实现。最后还是tl.store写入结果数组。行内分块reducedeftriton_row_sum(x:torch.Tensor,BLOCK_SIZE:int1024)-torch.Tensor:M,Nx.shape ytorch.empty(M,devicex.device,dtypex.dtype)row_sum_kernel[(M,)](x,y,N,BLOCK_SIZEBLOCK_SIZE)returny triton.jit defrow_sum_kernel(x_ptr,out_ptr,N,BLOCK_SIZE:tl.constexpr):rowtl.program_id(0)#One row:T1 T2 T3 T4|T1 T2 T3 T4|T1 T2 T3 T4acctl.zeros([BLOCK_SIZE],dtypetl.float32)forstart inrange(0,N,BLOCK_SIZE):colsstarttl.arange(0,BLOCK_SIZE)maskcolsN xtl.load(x_ptrrow*Ncols,maskmask,other0.0)accx resulttl.sum(acc,axis0)tl.store(out_ptrrow,result)softmax的核心是两次规约得到max和sum。这两部分刚才都是一个block就做完了。这里的block实际上就对应CUDA里的block每个block把数据从全局内存加载到共享内存再执行计算问题在于共享内存是有限的因此很多时候列数很大的时候不能让一整行都一次load全部加载应该在分成多次加载。这里给出一个reduce sum的行内分tile做法block大小固定为1024列数为N的话需要加载N/1024向上取整次每次把1024个元素load到共享内存进行计算。对于sum操作来说分块求和再把每一块的和加起来是等价的。Matmuldeftriton_matmul_relu(a:torch.Tensor,b:torch.Tensor):assert a.is_cuda and b.is_cuda assert a.is_contiguous()and b.is_contiguous()assert a.shape[1]b.shape[0]M,Ka.shape K,Nb.shape ctorch.empty((M,N),devicea.device)BLOCK_M,BLOCK_N,BLOCK_K64,64,32grid(triton.cdiv(M,BLOCK_M),triton.cdiv(N,BLOCK_N))matmul_relu_kernel[grid](a,b,c,M,N,K,a.stride(0),a.stride(1),b.stride(0),b.stride(1),c.stride(0),c.stride(1),BLOCK_M,BLOCK_N,BLOCK_K,)returnc triton.jit defmatmul_relu_kernel(a_ptr,b_ptr,c_ptr,M,N,K,stride_am,stride_ak,stride_bk,stride_bn,stride_cm,stride_cn,BLOCK_M:tl.constexpr,BLOCK_N:tl.constexpr,BLOCK_K:tl.constexpr,):pid_mtl.program_id(0)pid_ntl.program_id(1)indices_mpid_m*BLOCK_Mtl.arange(0,BLOCK_M)indices_npid_n*BLOCK_Ntl.arange(0,BLOCK_N)indices_ktl.arange(0,BLOCK_K)a_ptrs(a_ptrindices_m[:,None]*stride_amindices_k[None,:]*stride_ak)b_ptrs(b_ptrindices_k[:,None]*stride_bkindices_n[None,:]*stride_bn)acctl.zeros([BLOCK_M,BLOCK_N],dtypetl.float32)fork inrange(0,K,BLOCK_K):atl.load(a_ptrs,mask(indices_m[:,None]M)(indices_k[None,:]kK),other0.0,)btl.load(b_ptrs,mask(indices_k[:,None]kK)(indices_n[None,:]N),other0.0,)acctl.dot(a,b)a_ptrsBLOCK_K*stride_ak b_ptrsBLOCK_K*stride_bk acctl.maximum(acc,0.0)c_ptrs(c_ptrindices_m[:,None]*stride_cmindices_n[None,:]*stride_cn)tl.store(c_ptrs,acc,mask(indices_m[:,None]M)(indices_n[None,:]N),)这里用的就是第五讲讲过的矩阵乘法tiling策略如上图对结果数组C分tile每个(d,d)的c tile会需要(d,k)的A tile,(k,N)的B tile。如果NM很大贡献内存会放不下这就用到我们前面reduce算子提到的技巧block内用一个循环继续分块每次只load一小块到共享内存里。对于这一小块利用tl.dot(a, b)直接计算这是以block为基本单位的好处不用再操作线程具体做三重循环的矩阵乘法也就是我们不需要做绿色部分的操作。对于每个小块的matmul结果都累加到c tile临时数组arr最后把acc写回c数组
阅读完成 · 觉得有帮助?
咨询建站