从 AI 到算子这中间的距离比很多人想的要短。你用 PyTorch 写模型的时候一个torch.matmul、一个F.layer_norm背后其实是框架把整张神经网络先翻译成计算图再让一个个算子去执行。所谓算子就是最基础的计算单元卷积、矩阵乘、ReLU、Softmax都算算子。框架最终会把这些算子丢给底层的算子库比如 GPU 上的 cuDNN、CPU 上的 oneDNN或者各类 AI 加速卡自带的算子库由它们在硬件里完成并行计算。那为什么还需要你自己来写算子因为库提供的算子往往不是为你当前的模型结构量身定制的。有一类很典型的场景你想把多个算子的计算合并成一个减少中间结果的显存读写或者你正在跑一篇新论文里的新激活函数库里根本没有现成实现又或者你要在定制硬件上支持某种低精度格式官方算子来不及更新。这些都是算子开发者真正要解决的问题。这篇文章面向已经会写点 PyTorch 的同学从 AI 侧的直觉讲起带你把第一个算子写出来、跑起来、再接回训练流程再往后讲一点融合算子的做法和调优坑。1. AI 与算子先搞清楚我们到底在做什么1.1 为什么神经网络的计算会被拆成算子神经网络训练和推理的核心就是大量张量计算。PyTorch 这类框架不会直接逐行执行你写的y x bias这种 Python 表达式而是把模型结构转换成一个有向无环的计算图图里的节点就是一个算子边就是节点间的张量。训练时反向传播需要的梯度计算也会附加到这张图上每个算子同时提供 forward 和 backward 实现。这样做最大的好处是框架可以在运行时根据硬件环境、输入形状、精度要求来动态选择具体实现。比如torch.add可以选一个简单 kernel也可以用向量化指令torch.matmul在 A100 上可能会自动切分到 Tensor Core在 CPU 上则走 oneDNN 的 JIT 路径。算子作为统一调度单元把上层模型语义和下层硬件能力解耦了。但这也带来一个副作用同一个算子在不同模型里可能是鸡肋。你要在自回归生成里跑matmul mask softmax如果让框架分别调度三个算子中间张量要在显存里写一遍再读一遍带宽和时间都浪费了。于是你会发现很多新模型的高性能实现本质上都依赖融合算子和自定义算子这两板斧。1.2 哪些场景需要真正的算子开发不是所有 AI 工程师都需要写算子但下面几类场景手写算子几乎是绕不开的路。第一种是研究型新算子。注意力机制刚火那会儿flash_attn这名字还没人知道要高效实现它就得自己写融合 kernel。你做一篇新论文里的激活函数PyTorch 里没有原生的mamba_ssm或者glu也必须先写一个能跑通版本的算子。第二种是算子融合。像 LayerNorm、Softmax、GELU 这类元素级和归约级操作混合的模块在 Transformer 里经常出现。堆叠官方算子虽然能跑但中间张量的读写在显存带宽上非常吃亏。经验数据是一个简单模块的耗时里kernel 启动和中间张量拷贝往往占一半以上融合之后经常能省掉 30%~50% 的耗时。第三种是新硬件、新数据类型。int8、fp8、bfloat16 刚进入主流框架时官方算子库不一定覆盖所有 shape。要在特殊硬件上跑模型或者把卷积、矩阵乘替换成稀疏版本也都需要自己开发算子。第四种是极致性能。大模型服务对推理延迟极其敏感很多公司会针对自家模型手写 decode 阶段的算子目的是把带宽打满、减少内存格兰。这个阶段写算子就不是会不会的问题而是必须会的工程能力。我自己的判断方法是先用 Nsight 或 PyTorch Profiler 看一遍模型如果发现很多短 kernel 之间频繁发生大张量读写GPU 利用率不高那就有算子开发的优化空间。如果没有这种问题习惯性手写算子往往吃力不讨好。2. 算子开发的硬件与软件栈站在哪里下手2.1 从 AI 芯片到指令GPU 与 AI 加速卡的并行模型一个算子最终要跑在硬件上你得先理解硬件长什么样。GPU 的核心特点是大量线程并行执行同样的指令这就是 SIMT 模型。GPU 里有成百上千个计算核心每 32 个线程组成一个 warp硬件以 warp 为单位做调度。如果一个 warp 里的线程都在做同样的加法那一条指令可以同时作用于 32 份数据效率极高。反过来如果线程之间分支严重比如if (id % 2) ... else ...那硬件就得分别执行两边性能会掉不少。AI 加速卡NPU/DSA和 GPU 又不同。很多加速卡内部有专门的矩阵乘法单元、向量单元和标量单元它们共享一套存储和搬运机制开发算子时更像是编排数据搬运和三种计算单元之间的流水。以昇腾 AI 处理器的 Ascend C 开发接口为例它把算子抽象成矢量编程模型让你在高层描述输入、计算、输出编译器再去映射到硬件流水上而不需要你直接手写汇编。但底层要关注的仍是数据在什么地方、什么时候搬运、计算单元是否在等待数据。CPU 的并行度最低单核频率和整数性能最高对分支和复杂控制流的容忍度也高。写 CPU 算子的关键是缓存命中尽可能把数据停留在 L1/L2减少内存访问。2.2 开发选型CUDA、Ascend C、TVM / Triton选择哪套开发栈取决于你的目标硬件和控制的粒度。开发栈适合场景控制粒度学习成本CUDA CNVIDIA GPU 高性能开发线程、共享内存、PTX 都可控高Ascend C昇腾 AI 处理器的算子开发矢量、矩阵、标量流水可编排高Triton快速手写 GPU kernel编译器生成大部分细节中等TVM多硬件后端算法调度表达式 schedule高oneDNN/cuDNN用现成算子不可控低如果是刚入门我建议先选一套能动手的环境CUDA 生态最通用调试资料也最多。这篇就以 CUDA C 为例做演示但其中“线程组织、访存合并、归约、融合”的思想完全可以平移到你实际使用的加速卡上。换到 Ascend C 时你只是把 CUDA 的 grid/block 概念映射成它的并行编程模型把 shared memory 换成它的统一缓存核心思路没变。3. 实操从零写一个向量加法算子3.1 开发环境准备先确认你手里有什么写 CUDA 算子不需要特别复杂的工程环境但你得先知道自己的 GPU 型号和工具链版本。第一步是运行nvidia-smi看一下显卡名称、驱动版本、显存大小。然后确认 CUDA 编译器是否存在运行nvcc --version。如果你的环境里没有 nvcc只有驱动那说明只装了 runtime需要补一个 CUDA Toolkit。一个容易被坑的点是nvidia-smi显示的 CUDA Version 是驱动支持的版本上限不一定是本机 nvcc 版本。装 PyTorch 的时候会自带 CUDA runtime但用 nvcc 编译扩展时系统 PATH 里的 nvcc 可能来自另一个版本这会导致编译出的 kernel 与 PyTorch 的 CUDA runtime 不匹配运行时出现 no kernel image 或加载失败。遇到这种问题最简单的方法是统一到 PyTorch 配套的 CUDA 版本。我建议你在 Python 里先跑一句import torch print(torch.version.cuda) print(torch.cuda.get_device_name(0))确认 PyTorch 能正常调用 GPU然后再写算子。无论是从零编译还是用 PyTorch 的 extension 机制都能少踩很多环境坑。3.2 实现 kernel一个可以编译运行的向量加法最简单的算子就是向量加法C A B。这个例子的价值不是算法本身而是让你看清 GPU kernel 的代码结构、线程索引怎么写、怎么处理数组边界。#include cuda_runtime.h #include cstdio #include cstdlib __global__ void vector_add(const float* A, const float* B, float* C, int n) { int idx blockIdx.x * blockDim.x threadIdx.x; if (idx n) { C[idx] A[idx] B[idx]; } }这个 kernel 里blockIdx.x、blockDim.x、threadIdx.x分别表示当前线程所在的块编号、每块的线程数、块内线程编号。三者组合起来就是全局线程编号。if (idx n)用来处理末尾不足一个 block 的余数线程。host 端完成内存分配、数据拷贝、kernel 启动和结果拷回int main() { const int n 1024 * 1024; size_t bytes n * sizeof(float); float *hA (float*)malloc(bytes); float *hB (float*)malloc(bytes); float *hC (float*)malloc(bytes); for (int i 0; i n; i) { hA[i] i * 1.0f; hB[i] (float)i; } float *dA, *dB, *dC; cudaMalloc(dA, bytes); cudaMalloc(dB, bytes); cudaMalloc(dC, bytes); cudaMemcpy(dA, hA, bytes, cudaMemcpyHostToDevice); cudaMemcpy(dB, hB, bytes, cudaMemcpyHostToDevice); int blockSize 256; int gridSize (n blockSize - 1) / blockSize; vector_addgridSize, blockSize(dA, dB, dC, n); cudaDeviceSynchronize(); cudaMemcpy(hC, dC, bytes, cudaMemcpyDeviceToHost); // 校验若干位置再释放内存 return 0; }编译命令nvcc -O2 -archsm_80 vector_add.cu -o vector_add ./vector_add-arch后面是你 GPU 的 compute capability。比如 A100 是sm_804090 是sm_89H100 是sm_90。如果你不确定可以先省略-arch让 nvcc 生成兼容性更高的 PTX再由驱动 JIT 编译但那样性能未必最优真正跑模型时还是应该指定目标架构。3.3 把它接回 PyTorch让自定义算子和框架共用一个显存单机跑的示例代码跟实际 AI 工程之间还差一个“接回框架”的环节。你希望算子接收 PyTorch Tensor直接在 GPU 显存上计算并且参与自动求导。PyTorch 提供了torch.utils.cpp_extension.load可以在 Python 里直接编译 C/CUDA 扩展。准备两个文件一个.cpp用来绑定函数一个.cu用来放 kernel。.cpp的绑定大致长这样#include torch/extension.h #include cuda_runtime.h void vector_add_forward(torch::Tensor A, torch::Tensor B, torch::Tensor C) { int n A.numel(); int blockSize 256; int gridSize (n blockSize - 1) / blockSize; vector_add_kernelgridSize, blockSize( A.data_ptrfloat(), B.data_ptrfloat(), C.data_ptrfloat(), n); } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def(forward, vector_add_forward, vector add forward); }.cu里放的是前面写的 kernel。然后在 Python 里编译加载from torch.utils.cpp_extension import load vector_add load( namevector_add_cuda, sources[vector_add.cpp, vector_add_kernel.cu], extra_cuda_cflags[-O3, -archsm_80], verboseTrue )写完就可以测试import torch a torch.randn(1024, 1024, devicecuda) b torch.randn(1024, 1024, devicecuda) c vector_add.forward(a, b) torch.testing.assert_close(c, a b)这里有个很容易忽略的事传入的 Tensor 必须是 contiguous。PyTorch 里a.T会得到非连续存储的张量如果直接把data_ptr交给 kernel计算就会错。稳妥做法是在绑定函数里调用.contiguous()或提前在外面转好。4. 核心细节解析索引、访存与并行4.1 GPU 线程模型Grid、Block、Thread 怎么帮助并行写 CUDA 算子本质上是在组织工人干活。整个 GPU 上有大量计算核心你通过gridSize, blockSize一次性声明要启动多少个线程块、每个块里有多少线程。线程执行时由硬件自动把它们分成 warp每个 warp 32 个线程协同调度。一个新手最容易犯的错误是把threadIdx.x当成全局唯一的线程编号。其实threadIdx.x只是块内编号多个 block 里都有编号为 0 的线程。因此你需要用blockIdx.x * blockDim.x threadIdx.x算出全局下标。如果只用threadIdx.x去访问数组所有 block 都会访问同一段数据结果自然是错的。blockSize 的选择也影响性能。太小单个 SM 上能同时寄存的线程少无法藏住访存延迟太大每个线程能用的寄存器变少占用率也不一定更好。实际开发中我把 128、256、512 各跑一遍看谁的带宽最高然后选那个值。向量加法对 blockSize 不敏感但 reduce 类算子影响很大后面会再说。4.2 合并访存与 shared memory你的算子快不快的分水岭GPU 访存指令是按 warp 批量处理的一个 warp 的 32 个线程同时发起访存时如果它们访问的地址是连续的一段硬件可以合并成少数几次事务这就是合并访存。一旦线程访问的地址跨步比如array[idx * 2]硬件就可能要拆成更多次事务带宽利用率直线下降。看一个简单的对比同样计算C[i] A[i] B[i]如果写成C[i * 2] A[i * 2] B[i * 2]那等于只有一半的线程访问有效数据另一半被浪费或者每个 warp 的访存事务翻倍。这不是算法错误而是访存模式上的性能坑。shared memory 是 GPU 里可以由开发者直接控制的片上存储速度远快于全局显存。它也有一个经典坑叫 bank conflict。简单说shared memory 被分为 32 个 bank硬件按地址均匀分布到不同 bank。如果同一个 warp 的多个线程同时访问同一个 bank但地址不同访问会被串行化。一个常见的规避方式是在数组第二维加上一个 padding比如把float tile[32][32]改成float tile[32][33]让相邻行的列错开 bank。对刚入门的朋友我强烈建议先别碰太复杂的 shared memory 优化第一步是保证合并访存第二步是减少 kernel 数量。这两个动作带来的收益最直接。4.3 向量化指令与数学函数利用硬件特性榨干带宽GPU 上很多访存单元一次能读 128 位的数据相当于 4 个 float。如果你用标量方式每次只读一个 float硬件指令数量就会多出几倍。向量化写法就是把四个连续元素拼成一个float4一次加载、一次计算、一次存回。比如向量加法可以这样改__global__ void vector_add_vec4(const float4* A, const float4* B, float4* C, int n) { int idx blockIdx.x * blockDim.x threadIdx.x; if (idx n) { float4 a A[idx]; float4 b B[idx]; C[idx] make_float4(a.x b.x, a.y b.y, a.z b.z, a.w b.w); } }注意这里的n是 float4 元素的个数调用时需要把总元素数除以 4并单独处理剩余元素。实测在带宽受限的算子上向量化通常能带来 20%~50% 的提升代价是代码稍复杂。另外数学函数也要注意。GPU 上的expf、sqrtf如果你不追求极致精度可以用__expf、__frsqrt_rn等自带近似版本或者在编译器开-use_fast_math。但工程上要小心近似函数可能让训练结果产生微小差异推理时在意绝对精度的层比如 LayerNorm 的 rsqrt要谨慎使用。5. 进阶从基础算子到融合算子LayerNorm 融合5.1 为什么要做融合kernel 启动开销与中间张量读写向量加法只是热身真正体现算子开发价值的是融合。Transformer 里最常见的 LayerNorm按官方算子写的话至少有三个 kernel一个是求每行的均值与方差一个是做归一化如果还要接一个残差加又多一个 kernel。每个 kernel 启动都有微秒级开销更重要的是每个 kernel 都要把 X 从全局显存读一遍、把中间结果写一遍。一个很直接的经验是算子如果计算强度低也就是算得少、访存多那么决定耗时的主要是数据搬移量而不是浮点运算量。LayerNorm 每读一个 float 无非是加几次、乘几次内存带宽就是瓶颈。所以把均值、方差、归一化、乘 gamma、加 beta 全部揉进一个 kernelX 只在主显存里读一次Y 只写一次中间变量完全留在寄存器和 shared memory 里收益非常可观。我自己在调这类融合算子时惯用的判断标准是算术强度计算量 / 访存量。如果这个比值小于 1基本就是访存瓶颈融合是主要优化手段。5.2 一个 LayerNorm 融合算子的设计假设输入 X 的形状是[rows, cols]每行需要独立归一化。最简单的并行方案是每个 block 负责一行块内所有线程协作完成一行的归约和归一化。先给出公式mean sum(x) / cols var sum((x - mean)^2) / cols y (x - mean) / sqrt(var eps) * gamma beta朴素写法是两遍第一遍算 mean第二遍算 var。但这样要读两次 X。更稳的方式是单遍 Welford 算法在线更新 mean 和 M2一次循环里就能累积出方差信息。代码如下__global__ void layernorm_fused_kernel(const float* __restrict__ X, const float* __restrict__ gamma, const float* __restrict__ beta, float* __restrict__ Y, float eps, int cols) { int row blockIdx.x; int tid threadIdx.x; int stride blockDim.x; const float* x_row X row * cols; float mean 0.f; float M2 0.f; int count 0; for (int col tid; col cols; col stride) { float val x_row[col]; count 1; float delta val - mean; mean delta / count; float delta2 val - mean; M2 delta * delta2; } // 这里需要做块内归约把每个线程的 mean 合并为全局 mean_final // 再把每个线程的 M2 合并为全局 M2_finalWelford 归约比 sumSq 更稳 // 可以使用 warp shuffle shared memory 实现或依赖 cub::BlockReduce float var M2_final / count; float inv_std rsqrtf(var eps); for (int col tid; col cols; col stride) { float y (x_row[col] - mean_final) * inv_std; Y[row * cols col] y * gamma[col] beta[col]; } }上面说的块内归约是绕不开的核心环节。具体做法是先把每个线程处理出的局部 mean/M2 放到 warp 级 reduction用__shfl_xor_sync把 32 个线程的值两两相加然后每个 warp 选一个线程把结果写进__shared__最后由第一个 warp 做最后一次归约得到整行的值。这个流程是所有 reduce 类算子的基本功值得单独练一遍。一个更简单的替代方案是直接求sum和sumSq最后var sumSq / cols - mean * mean。好处是归约好写、代码短但数值稳定性差一点在 fp16 或某些极端分布下可能算出负方差。如果只是学习可以先从 sum/sumSq 版本开始再升级到 Welford。5.3 手工优化清单从能跑变成能打写出正确融合 kernel 只完成了一半。下面这份清单是我实际调算子时基本都会过一遍的每个 block 处理多行而不是一行。小cols时一行的归约不足以占满整个 block让每个 block 负责连续多行减少 block 总数提高占用率。动态 shared memory 配合extern __shared__使用避免为一个形状写死缓冲区。向量化读取 X。如果能保证cols是 4 的倍数用float4把四列一次读入。给指针加上__restrict__帮助编译器判断 X、gamma、beta 之间没有别名从而更好地优化循环。归一化阶段的gamma[col]和beta[col]尽量走只读缓存可以用__ldg或确保它们是 const 限定。起始 blockSize 不要拍脑袋定先测 128、256、512、1024观察哪个配置下带宽利用率最高。这些优化没有一条是高深理论但每一条都可能带来几个百分点的性能提升叠加起来差距就拉开了。6. 常见问题与排查技巧实录6.1 编译与链接阶段的坑第一次编译 CUDA 算子最常见的错误就是nvcc报 unsupported gpu architecture。原因很简单你没指定-arch或者指定了和本机 GPU 不匹配的架构号。解决办法是查清楚 compute capability比如nvcc -O3 -gencode archcompute_80,codesm_80 vector_add.cu另一个坑是链接错误在同一份工程里.cpp和.cu都定义了同名全局函数。如果 CUDA kernel 需要被 C 函数调用记得在.cu里用extern C声明接口或者在头文件里做好声明否则链接阶段会抱怨找不到符号。还有一类问题跟 PyTorch extension 加载相关编译成功了但 import 时提示 CUDA driver version insufficient或者加载进来的模块里 kernel 无法执行。这通常是你本机 nvcc 版本与 PyTorch 预编译的 CUDA runtime 版本不一致。我一般用torch.utils.cpp_extension传入的extra_include_paths对准 PyTorch 自带的 CUDA 头文件而不是系统路径里的版本。6.2 运行期 illegal memory access如何快速定位越界illegal memory access属于运行时错误kernel 已经启动但你访问了非法地址。最直接的办法是跑compute-sanitizercompute-sanitizer --tool memcheck ./vector_add如果使用 PyTorch extension也可以在 Python 里用同样的工具包住整个测试脚本。它会准确报告是哪个 kernel、哪个线程、访问了什么地址省去一行行打印的苦功夫。导致越界最常见的原因就是边界没处理好。比如向量长度 n 不是 blockSize 的整数倍又没有写if (idx n)线程就会访问数组末尾之外。处理办法是保留边界判断但为了性能可以让主循环用 grid-stride loop把整段数据先按对齐处理剩余尾部再慢速处理。6.3 性能不达预期别只看 kernel 时间一个很常见的心态是写完自定义算子跑一下发现比torch.add慢于是怀疑是自己写错了。其实不一定。torch.add背后是高度调优的模板库本身可能已经向量化且选好了最佳 block 配置你手写一个朴素 kernel 打不过它是正常的。更合理的流程是先用 Nsight Systems 看整个模型的 kernel 时间分布再用 Nsight Compute 看单个 kernel 的带宽利用率、计算利用率和 warp 占整率。如果你的算子把访存带宽跑到了 80% 以上就算打不过官方库也已经算合格了。如果利用率只有 30%那就去检查合并访存、向量化和 blockSize这三项优化空间最大。另外测性能时不要只测一个超大 shape。LayerNorm 的cols从 256 到 8192 变化时最佳 block 配置和策略完全不同。多测几组 shape把数据记录下来再选方案。最后分享一个我自己的习惯手写算子接回 PyTorch 后我会刻意加上一个torch.testing.assert_close的断言把rtol设到 1e-4 左右每次跑模型前先通过正确性检查。算子开发最怕的不是慢而是结果悄悄错掉一旦浮点误差累积到整个模型里排查起来比性能问题痛苦得多。先保证正确再谈优化这句话值得每个刚入门的人贴在显示器旁边。
