ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

从GEMM到Transformer:手写CUDA算子优化的完整路线与性能调优实践

2026/9/7 15:06:32 拓冰建站 浏览量
从GEMM到Transformer:手写CUDA算子优化的完整路线与性能调优实践 5. 写在Week4末尾的体会终于把这一周的东西整理完了。说句实话这周是我接触CUDA以来最“烧脑”但也最“上瘾”的一周——从最简单的GEMM开始一路优化到能自己手写Transformer里几个关键算子的kernel最后还能用Nsight Compute给性能瓶颈“定罪”。整个过程很像在做一道需要反复打磨的算法题只不过这台“电脑”变成了带几千个核心的GPU。这篇总结适合几类人看正在系统学CUDA、想搞懂深度学习算子优化的人被各种CUDA环境问题折磨到怀疑人生的人以及准备面试推理引擎、算子优化岗位的同学。我不会上来就贴一堆高深的优化技巧而是尽量把“为什么要这么做”“踩过的坑长什么样”都讲清楚。文章里所有的数值和现象都是我这一周在自己机器上实测出来的配置是RTX 3090 Ubuntu 22.04 WSL2PyTorch 2.8.0 CUDA 12.1组合包。你机器不同没关系方法论完全通用。1. 为什么Week4先死磕GEMM1.1 GEMM是深度学习的“基本盘”很多人一上来就想写FlashAttention那种炫酷的kernel我真心不建议。你在Transformer里看到的绝大多数计算本质上都能归约到GEMM全连接层就是GEMM卷积通过im2col或者隐式GEMM实现Attention里的QK^T和scoreV是GEMMFFN的两层线性变换是GEMM包括多卡张量并行里的矩阵分块底层也是GEMM的拆分逻辑。所以GEMM优化是所有算子优化的“基本盘”。把GEMM的共享内存、bank conflict、向量化、双缓冲这套方法论吃透再去看FlashAttention的tiling策略你会觉得特别熟悉因为它的核心就是在处理QK^T和PV这两个大GEMM中间夹了一个softmax。这周我给自己定的目标是把一个朴素的GEMM kernel从不到200 GFLOPS的实测性能优化到接近8 TFLOPS。数字本身不重要重要的是我真正理解了每一步优化到底在解决什么问题。1.2 我的GEMM优化路线图Naive、Tiling、向量化、双缓冲先说下朴素版本长什么样。假设我们要算 C[M,N] A[M,K] B[K,N]最直接的想法是每个线程算一个输出元素__global__ void gemm_naive(const float* A, const float* B, float* C, int M, int N, int K) { int row blockIdx.y * blockDim.y threadIdx.y; int col blockIdx.x * blockDim.x threadIdx.x; if (row M col N) { float sum 0.0f; for (int k 0; k K; k) { sum A[row * K k] * B[k * N col]; } C[row * N col] sum; } }这段代码逻辑完全正确但性能惨不忍睹。为什么每个线程要读A的一整行和B的一整列block里的线程之间没有任何数据复用。比如16x16的block16行A和16列B其实总共只涉及161632段数据但每个线程各自从全局内存取重复读了16倍以上。全局内存带宽就那么大时间全部浪费在等数据上。第一步肯定是共享内存tiling。思想很直白一个block先把需要用到的A子块和B子块搬进共享内存然后block内的线程反复从共享内存取数计算。共享内存比全局内存快一个量级关键是它让数据被“缓存”在了离计算最近的地方。实现方式也不复杂——定义一个block负责计算一块 TILE_M x TILE_N 的输出比如32x32。每个线程仍然算一个输出元素但在算之前整个block协作把A的一个32x32 tile和B的一个32x32 tile加载到共享内存#define TILE 32 __shared__ float As[TILE][TILE]; __shared__ float Bs[TILE][TILE]; int tx threadIdx.x, ty threadIdx.y; int row blockIdx.y * TILE ty; int col blockIdx.x * TILE tx; float sum 0.0f; for (int k0 0; k0 K; k0 TILE) { As[ty][tx] A[row * K k0 tx]; Bs[ty][tx] B[(k0 ty) * N col]; __syncthreads(); for (int k 0; k TILE; k) { sum As[ty][k] * Bs[k][tx]; } __syncthreads(); } C[row * N col] sum;注意这里的两次__syncthreads()第一次是确保共享内存里的数据写完了再开始读取计算第二次是确保所有线程算完了再覆盖下一轮数据。没有这一步会出现典型的数据竞争算出来的结果在边界处时对时错而且很难复现是最恶心的bug之一。优化到这一步性能大概能从200 GFLOPS涨到1-2 TFLOPS但离硬件上限还差得远。接下来的关键操作是向量化加载把float换成float4一次读16字节。这对全局内存和共享内存都有好处但要求地址16字节对齐所以通常在K维度上确保长度是4的倍数。做的时候有个小坑shared memory声明时最好用float4数组而不是float数组否则编译器很难自动向量化。双缓冲共享内存的加载是有延迟的如果一整个block在那里傻等数据算完再加载下一块计算单元就空转了。双缓冲的思路是在计算当前tile的时候提前把下一个tile的数据加载到另一块共享内存。在Ampere架构上可以用cp.async直接发出异步拷贝不占用寄存器也不阻塞线程这一招对隐藏内存延迟效果极其明显。每个线程算多个输出把thread coarsening加上例如每个线程算4x416个输出元素。这样做的好处是增加数据复用率同一个A子块的数据可以被算4次同时减少block数量调度开销也更低。这一套组合拳打下来我的GEMM实测性能从不到200 GFLOPS提升到了8 TFLOPS左右。虽然离3090的FP32峰值35 TFLOPS还有差距但从“明显错误”到“能看”这个跨越是巨大的。1.3 Bank Conflict与Padding一个细节拖垮整个kernel如果只看上面那段tiling代码你可能觉得单核利用率已经不错了。但我在实测时发现SM的利用率经常只有20%上下后来用Nsight Compute一查罪魁祸首是共享内存bank conflict。共享内存的硬件结构是把连续的存储空间切分成32个bank每个bank每个时钟周期只能处理一个访问请求。一个warp有32个线程如果它们同时访问同一个bank里的不同地址硬件就必须把这些访问拆成多个周期串行执行这就叫bank conflict。等于说共享内存带宽被打折了。最常见的冲突场景就在矩阵转置或者按列访问矩阵元素时。比如Bs[k][tx]这种访问方式在B以行主序存储时warp里的32个线程访问的是同一行的不同列——也就是第32列、第33列、第34列……它们落在同一个bank里于是一个本来一个周期就能完成的load变成了32个周期。解决办法特别简单粗暴给共享内存数组的每一行加一个float的padding。把Bs[TILE][TILE]改成Bs[TILE][TILE 1]。这样第i行的第j个元素和真实内存地址之间错开了一个floatwarp访问同一行的不同列时索引自然分散到不同bank冲突就消失了。这个细节让我明白一个道理算力虽然重要但访存路径上任何一个“看似无伤大雅”的设计都可能让kernel性能掉一个数量级。所以后续每次优化我都拿着Nvidia的Nsight Compute看一下shared__st_bank_conflicts、shared__ld_bank_conflicts两个指标比盲调快太多。2. 从GEMM到Transformer把注意力机制拆成算子清单2.1 Transformer的全部关键算子拆解有了GEMM的底子再看Transformer就轻松很多。我做的第一件事不是写代码而是把整个Transformer前向计算拆成一张算子清单然后按“计算密集”和“访存密集”分类Embedding查表Gather操作访存密集GPU上其实不太划算更多是内存带宽受限位置编码PE逐元素sin/cos访存密集计算量相对小Q/K/V投影三个GEMM实际工程上通常拼成一个大的GEMM做计算密集Attention ScoreQ与K^T的GEMM再加scale和mask计算密集Softmax逐行归一化访存密集是Attention里最容易忽略的瓶颈Attention Outputscore与V的GEMM计算密集Output Projection又一个GEMMLayerNorm先求均值方差再归一化访存密集FFN两个GEMM夹一个GELU激活第一个GEMM计算量尤其大因为中间维度通常是4倍。按时间占比看在Decoder里GEMM类算子能占到70%到90%的耗时Softmax、LayerNorm、Embedding加起来占比不高但它们在端到端推理里会造成频繁的kernel launch和显存读写不能被忽视。所以我在Week4后半段把重点放在两件事上手写一个融合的Softmax Attention kernel以及把位置编码PE在GPU上的正确实现方式搞明白。2.2 手写一个融合的Softmax Attention Kernel先说说朴素的Attention实现方式PyTorch里直接写softmax(Q K^T / sqrt(d)) V这是最直观的写法但性能上有一个很大的问题——QK^T的结果会被写回全局内存softmax再读一遍然后乘V再写一遍。这一来一回中间矩阵[seq_len, seq_len]在HBM和SM之间反复横跳如果seq_len是4096光这个中间矩阵就要占64MB非常浪费。优化思路有两个方向。第一把softmax融合到矩阵乘法里减少中间矩阵的读写第二把QK^T、softmax、V三者合到一个kernel里这其实就是FlashAttention的雏形。我先说softmax本身怎么在kernel里写。朴素softmax需要两遍扫描第一遍找最大值第二遍算exp和归一化。但开两个kernel显然不划算。在线softmaxonline softmax可以只扫一遍假设当前已经处理到第t个元素维护当前最大值m_t和累积指数和l_t。读入新元素x后float m_new fmaxf(m_old, x); l_new l_old * expf(m_old - m_new) expf(x - m_new);这样一遍扫描就能拿到真实的max和sum。之后再从头读一遍算最终的exp(x - m_max) / l或者如果后续还要做矩阵乘V就边更新边累加。这个“边扫描边累积”的思路看起来简单但它是FlashAttention能够节省访存的核心没有它每次都要多读一遍输入。我在自己写的Attention kernel里做了一件更激进的事把QK^T的scale和mask也融合进去。每个block负责query子块和key子块的一小段点积算完直接乘上scale如果某个位置需要mask比如padding就直接把score设成-inf这样exp之后天然是0不需要额外的分支。然后在线softmax维护running max和sum同时累加V的输出贡献。整个过程中间不会生成一份完整的[seq_len, seq_len]矩阵显存占用低很多。写这个kernel的时候有两个我特别想强调的坑一是数值稳定性。如果不减去max直接算exp(x)当score很大时expf会溢出变成infsoftmax结果直接出NaN。这也是为什么在线softmax里必须维护running max的原因。二是__syncthreads()的摆放位置。因为softmax的归一化涉及整个row的统计量而每个row可能跨多个block协作计算block之间需要同步或者干脆让一个block处理一整行所以要么限制每个block处理的行数要么用grid.sync之类的机制。我前期踩的坑是不同block各算各的最后归一化因子对不上结果错得离谱。最终这个融合kernel在我机器上比“分开写”的版本快了大概2.5倍中间矩阵的显存占用直接省掉了。这个优化带来的爽感比GEMM调优还强因为你能感觉到自己的“算法意识”在起作用而不只是堆优化技巧。2.3 Embedding与位置编码PE在GPU上的正确打开方式热词里很多人搜“transformer架构嵌入表示层pe计算”说明这块对新手确实容易乱。原版Transformer的位置编码公式长这样PE(pos, 2i) sin(pos / 10000^(2i/d_model)) PE(pos, 2i 1) cos(pos / 10000^(2i/d_model))其中pos是token在序列里的位置i是维度索引范围是0到d_model/2 - 1。这个编码的本质是给每个位置生成一个“频率不同”的正弦波让模型能够感知相对位置关系。在GPU上实现PE最稳妥的做法是预计算一个[seq_len, d_model]张量一次性拷到显存里每个batch重复用。别在每次前向的时候在GPU上现算一遍纯属浪费。用PyTorch预计算的代码几乎任何教程都会给import torch import math def positional_encoding(seq_len, d_model): pe torch.zeros(seq_len, d_model, dtypetorch.float32) position torch.arange(seq_len, dtypetorch.float32).unsqueeze(1) div_term torch.exp( torch.arange(0, d_model, 2, dtypetorch.float32) * (-math.log(10000.0) / d_model) ) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) return pe这段代码里有几个细节值得展开。第一div_term用的是exp(arange * (-log(10000)/d_model))而不是直接算10000^(2i/d_model)再取倒数。原因是前者可以预先算好一个长度为d_model/2的向量避免在GPU上每个位置都算一次pow(10000.0, ...)这既是数值稳定的写法也是性能友好的写法。第二一定要用float32来算不要用float16。虽然推理时模型权重可能是FP16但PE里有个1/10000^(2i/d_model)当i比较大时分母会变得巨大FP16很容易精度不够位置编码就是一堆噪声。如果真要手写一个CUDA kernel来生成PE也不复杂。每个线程负责一个(pos, dim)坐标通过dim % 2判断是算sin还是cos。但更高明的做法是把div_term作为常量数组传入kernel每个线程只需做一次position * div_term[dim / 2]和一次sin/cos调用把重复计算全部去掉。这个“常量前移”的思想在算子优化里非常通用几乎任何kernel都能用一遍。3. CUDA环境与工具链跑通算子的前提3.1 多版本CUDA共存PATH、LD_LIBRARY_PATH和软链接的坑这一周我的机器上至少同时出现过CUDA 11.8、12.1、12.4三个版本因为不同项目依赖的框架版本不一样。一开始我被环境问题折磨到想砸电脑后来才总结出一套相对稳妥的管理方式。先理解几个命令的区别。nvidia-smi右上角显示的是“当前驱动支持的最高CUDA版本”这并不等于你系统里实际安装的CUDA Toolkit版本。nvcc --version显示的才是Toolkit的版本。很多人一看nvidia-smi说支持12.1但nvcc --version还是11.8就以为环境坏了其实这完全正常驱动和Toolkit本来就是两码事。多版本共存时最忌讳的事情是手改/usr/local/cuda这个软链接。每次切换版本都去改软链接然后全局改PATH和LD_LIBRARY_PATH短期内能用但很容易把正在跑的服务搞挂而且切来切去迟早会忘记当前到底是哪个版本排查问题时异常痛苦。我现在的做法是每个项目一个conda环境在环境里装对应版本的PyTorch和CUDA相关依赖。完全不需要在系统层面做全局切换。如果一个项目必须用系统Toolkit编译CUDA扩展再单独指定CUDA_HOME和PATHexport CUDA_HOME/usr/local/cuda-12.1 export PATH$CUDA_HOME/bin:$PATH export LD_LIBRARY_PATH$CUDA_HOME/lib64:$LD_LIBRARY_PATH如果你管理的是多用户服务器想让所有用户默认都生效可以写到/etc/profile.d/cuda.sh里。但注意路径不要写死最好通过/usr/local/cuda这个软链接来指路这样后续版本升级时只需要更新软链接。额外提醒一个常见误区LD_LIBRARY_PATH里的lib顺序会影响运行时加载的cudart、cublas版本如果设置错了即使nvcc --version是对的运行程序时也可能加载到另一个版本的库。排查这类问题可以用ldd ./your_app看实际加载路径。3.2 WSL2安装CUDA的正确姿势现在我大部分调优工作都在WSL2里做说下正确姿势。Windows宿主机装好NVIDIA驱动后WSL2里直接运行nvidia-smi就可以看到GPU信息显示的还是和Windows相同的驱动版本。很多人误以为这样就算装好CUDA了结果一编译代码发现nvcc不存在——没错WSL2里的驱动由Windows共享但CUDA Toolkit需要自己在Linux里装。在WSL2里安装CUDA Toolkit去官网选择Linux、WSL-Ubuntu的安装包就行安装方式跟原生Ubuntu几乎一样。装的时候建议别用默认的/usr/local/cuda改软链接方式直接在WSL的~/.bashrc里配置export PATH/usr/local/cuda/bin:$PATH export LD_LIBRARY_PATH/usr/local/cuda/lib64:$LD_LIBRARY_PATHWSL2下的PyTorch组合我实测比较稳的组合是Python 3.10 PyTorch 2.8.0 CUDA 12.1。去PyTorch官网get-started页面选好配置复制它会给你生成的那条pip install命令别自己随便拼。pip install torch默认装的是CPU版本很多人装完发现torch.cuda.is_available()返回False十有八九是这里出了问题。安装完一定要跑这条命令验证python -c import torch; print(torch.__version__, torch.version.cuda, torch.cuda.is_available())如果输出True再跑torch.ones(1).cuda()确认实际能执行GPU计算。我在WSL2踩过最坑的一次是torch.cuda.is_available()返回True但一跑真实模型就报no kernel image问题出在PyTorch版本和GPU算力不匹配这个下面第四节详细说。3.3 用Nsight Compute确认优化方向优化kernel最忌讳“盲人摸象”。我见过太多人一上来就疯狂调block大小、换循环顺序结果性能纹丝不动因为没有定位到真正的瓶颈。Nsight Compute命令行工具ncu就是解决这个问题的。最基本用法ncu --set full ./your_app它会跑一遍程序然后给出一大堆性能指标。你不需要全看懂重点关注几个sm__throughput.avg.pct_of_peak_sustained_elapsedSM计算单元的利用率dram__throughput.avg.pct_of_peak_sustained_elapsed显存带宽利用率sm__warps_active.avg.pct_of_peak_sustained_active活跃warp数也就是occupancyl1tex__data_bank_conflicts_pipe_lsu_mem_shared_op_ld共享内存bank conflict计数。判断方法很简单如果dram__throughput接近80%以上而sm__throughput只有20%说明这个kernel是访存瓶颈memory-bound优化重点应该放在减少全局内存读写、增加数据复用、用向量化加载降低请求次数。反过来如果SM利用率高、DRAM利用率低说明是计算瓶颈compute-bound考虑减少冗余计算、提高指令级并行。我之前的GEMM kernel优化到1 TFLOPS左右就卡住了怎么调block大小都没用。上ncu一看dram__throughput已经98%SM只有11%根本原因就是朴素实现疯狂读全局内存。后续加上共享内存tiling后DRAM利用率一下子降到30%SM利用率升到60%以上性能自然就上来了。有一点要注意ncu在某些环境会报ERR_NVGPUCTRPERM权限错误一般用sudo ncu运行就能解决。如果你在容器里使用需要额外开性能计数器权限。这属于环境配置问题别在kernel代码上白费劲。4. 新手必看CUDA算子优化常见问题排查实录4.1 “no kernel image is available for execution”到底在说什么这个报错满屏都是尤其是在新显卡上跑老PyTorch的时候。完整报错通常是torch.acceleratorerror: cuda error: no kernel image is available for execution on the device它的真实含义是CUDA在加载某个kernel时找不到能在当前GPU上执行的“kernel image”。GPU的SASS指令是针对具体架构SM版本编译的比如RTX 3060是sm_86RTX 4060 Ti是sm_89。如果一个库编译时只包含了sm_80或sm_86的SASS那它在sm_89的卡上就找不到可执行代码。这就是为什么40系显卡用户经常遇到这个问题新卡架构太新老版本的PyTorch wheel里没有包含对应架构的内核。3060、4060 Ti这些卡如果你装的PyTorch版本过老也可能报同样错误因为老版本根本没有sm_86或sm_89的预编译kernel。解决办法按优先级排列升级PyTorch到与显卡架构匹配的新版本同时选择官方wheel对应的CUDA版本如果必须用旧PyTorch确保CUDA版本足够新例如CUDA 12.x的wheel通常覆盖更多新架构但这只是兜底自己编译PyTorch扩展时设置环境变量TORCH_CUDA_ARCH_LIST8.9;8.6或对应算力确保编译产物包含目标架构。排查之前先跑一下python -c import torch; print(torch.cuda.get_device_capability())确认当前GPU的算力。只有知道目标算力才能选对PyTorch版本。4.2 CUDA Samples找不到、nvidia-smi与nvcc版本不一致“CUDA Samples找不到”也是个高频问题。出现这个情况一般是安装Toolkit时没勾选Samples组件或者用的是精简安装。正常Samples应该出现在/usr/local/cuda/samples如果你没找到可以根据系统直接拉取对应版本的cuda-samples比如Ubuntu/Debian系apt-get install nvidia-cuda-samples装完的路径通常在/usr/share/nvidia-cuda-samples不在/usr/local/cuda下面别傻找。RedHat系用dnf install nvidia-cuda-samples类似。至于“nvidia-smi显示的CUDA版本和nvcc --version不一致”我在3.1节说过这是正常的。再强调一次nvidia-smi右上角是驱动支持的上限nvcc --version是你当前Toolkit版本。两者不一样不代表环境坏了只要满足“驱动版本 Toolkit所需版本”就行。驱动是向后兼容的新驱动可以运行旧Toolkit编译出来的程序。4.3 驱动、Toolkit、PyTorch、cuDNN的版本对应关系我把这四者的关系整理成一张速查表方便你排查问题时对照组件查看命令作用典型坑GPU驱动nvidia-smi底层的GPU驱动决定硬件可用性和CUDA运行时上限右上角版本不是Toolkit版本CUDA Toolkitnvcc --version编译CUDA代码链接库系统可以装多个版本需要管理PATHPyTorch内置CUDA Runtimetorch.version.cudaPyTorch运行时依赖的cudart/cublas与系统Toolkit不一致是正常的cuDNNcat /usr/local/cuda/include/cudnn_version.h深度学习卷积等算子的加速库需要和CUDA版本匹配这里有个重要的认知PyTorch的wheel包自己带了一套CUDA运行时依赖所以即使你系统里没装CUDA Toolkit也能跑import torch; torch.cuda.is_available()为True的模型推理。只有当你要从头编译CUDA扩展比如自己写kernel或者编译flash-attn这类库时才必须装完整的Toolkit。对于4060 Ti这类新卡很多教程会问你“4060ti支持的cuda版本”其实这是问错了方向。40系卡在CUDA 11.8和12.x下都能跑关键是PyTorch的wheel是否包含对应架构的SASS。版本选择应该以PyTorch官方支持矩阵为主而不是单独看显卡。4.4 我的排查顺序与工具清单这一周我几乎每天都要排查各种环境问题后来总结了一套固定顺序能解决90%的“莫名其妙”的错先确认GPU裸设备有没有被识别nvidia-smi。这一步挂了后面全白搭确认PyTorch能不能看到GPUpython -c import torch; print(torch.cuda.is_available())确认当前设备算力torch.cuda.get_device_capability()执行一个最小GPU算子torch.ones(1).cuda()如果需要编译扩展再确认nvcc --version和torch.version.cuda是否能对齐不对齐会有一堆链接报错怀疑环境变量有问题时用env | grep -i cuda和ldd看实际加载库路径。整个过程我基本不动系统级的配置所有依赖都锁在conda环境里。如果你发现自己经常要来回切换CUDA版本我给的建议也很朴素不要硬改全局环境用conda或者venv隔离每个项目一套环境。可能刚开始会觉得多占点磁盘但换来的是半年不用为环境问题掉头发这笔账怎么算都划算。5. 写在Week4末尾的体会这周最大的收获反而不是某个具体kernel提速了多少而是我终于建立了“性能问题要先量化再动手”的直觉。没上ncu之前我总凭感觉觉得瓶颈在某个循环花了一晚上重写结果性能纹丝不动上了profiler之后两分钟就发现DRAM带宽利用率满了算力根本没吃满。另一个体会是算子优化的核心永远在“访存”与“计算”的权衡。GEMM的tiling是在用共享内存换全局内存访问次数FlashAttention是在用重计算换HBM带宽PE预计算是在用显存换实时计算开销。想通这一层后再看任何一篇优化论文或者源码你都能猜到它大概在优化什么、为什么这样优化。最后分享一个我这周才养成的小习惯每次改kernel前先保存上一个版本的性能数据建议用一个简单的表格记下版本号、改动内容、SM利用率、DRAM利用率、耗时。改完一版就更新一行不要凭感觉。这个习惯帮我少走了很多回头路——很多时候你觉得新版本更快其实是测量误差有个基线数据对比就不会自欺欺人了。下周我打算沿着这条线继续深入FlashAttention的完整实现以及Grouped GEMM在多头注意力里的应用。如果你也在学CUDA算子优化或者刚被某个no kernel image报错逼疯欢迎直接在评论区聊聊你踩过的坑我也想知道大家还会遇到什么奇葩问题。