ARTICLE DETAIL

资讯详情

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

用CUDA自定义融合Softmax算子并封装进PyTorch

用CUDA自定义融合Softmax算子并封装进PyTorch 有一段时间我在训练一个 Transformer 模型profiling 跑下来令我意外的是占据大量时间的不是矩阵乘法反而是看起来不太起眼的 Softmax。查了一圈发现PyTorch 原生的 softmax 实现是通用版一个完整的 masked softmax 要被拆成乘法、加法、max、exp、sum、除法好几个 kernel 来回倒腾显存。那时候我就萌生了一个念头干脆自己写一个融合的 CUDA 算子。这篇文章就围绕这个话题来展开讲讲我如何用 CUDA 编程实现一个自定义 ScaledMaskSoftmax 算子并把它顺利封装进 PyTorch 的 autograd 体系里让训练代码里直接可以调用。如果你有 PyTorch 基础正准备接触自定义算子或者只是好奇 Attention 里的 softmax 究竟可以怎么优化这篇文章应该能给你一条清晰可复现的路径。1. 为什么要把一个“现成的Softmax”写成自定义CUDA算子1.1 通用算子路线的性能账本很多人一开始会有疑问torch.softmax不是已经很快了吗写自定义算子是不是有点多此一举我们先算一笔账。在标准的 Attention 计算里ScaledMaskSoftmax做的事情是对QK^T的结果除以sqrt(d_k)加上 mask然后做 softmax。如果你完全用 PyTorch 原生算子拼大概长这样scores torch.matmul(q, k.transpose(-2, -1)) scores scores / math.sqrt(d_k) scores scores mask probs torch.softmax(scores, dim-1)这四行代码看起来简洁但底层实际发生了什么每一步都是一个独立的 CUDA kernel数据在显存里的流动路径是这样的matmul 写完一份中间结果除法或乘法再读一次写一次加 mask 再读一次写一次softmax 内部为了数值稳定性还要先做一次 reduce max、一次 reduce sum最后做归一化。也就是说一个中间张量可能要被全局显存反复读写五六遍。如果你训练的是大 batch、大 seq_len 的模型这部分开销会被放大得非常明显。Attention 的分数矩阵是[B, H, S, S]当S1024、B8、H32时光这一个张量就是8*32*1024*1024*4字节大约 1GB 显存每多一次读写就是 2GB 的显存流量。所以通用实现虽然正确但绝对不是最划算的。1.2 ScaledMaskSoftmax到底长什么样在动手写代码之前先把目标算子定义清楚。输入是一个四维张量x形状是[B, H, S, S]表示 Attention 的原始分数。我们沿最后一维做 Softmax也就是针对每一个 query 位置对其所有的 key 位置做归一化。数学形式很简单score_ij x_ij * scale mask_ij output_ij exp(score_ij) / sum_j(exp(score_ij))这里的scale通常取1 / sqrt(d_k)。mask有两种常见用法一种是对 padding 位置加一个绝对值很大的负数让 softmax 之后概率接近 0另一种是因果掩码也就是j i的位置全部屏蔽用于自回归模型。我决定把这两种情况都放进同一个 kernel 里因为实际项目中往往不是只用一种把它们融合在一起调用起来才足够灵活。写自定义算子的核心收益就是把这些所有操作合并成一个 kernel从全局显存读一次分数在共享内存里完成 scale、mask、reduce、归一化最终写回一次结果。GPU 是典型的吞吐优先架构减少全局内存往返往往比减少计算量更能带来直观的提速。2. 从数学到线程映射Kernel设计先想清楚三件事2.1 softmax数值稳定性为什么非减max不可Softmax 的朴素实现是exp(x_i) / sum(exp(x_i))但exp的输入如果很大比如x_i100单精度浮点数直接溢出成inf。Attention 分数经过 scale 之后通常在一个可控范围内但加了 mask 或者训练初期参数不稳定时出现极端值是常有的事。所以数值稳定版的 softmax 一定会做一步变换m max_j(score_j) output_i exp(score_i - m) / sum_j(exp(score_j - m))减掉行最大值之后所有指数项的输入都不超过 0exp的结果落在(0, 1]区间彻底避免溢出。这个虽然是很基础的常识但写 CUDA kernel 的时候特别容易漏一旦漏掉可能在特定输入下产生nan而且这种 bug 非常难查。我的建议是不管你认为输入范围多安全稳定版变换一定要做。2.2 mask的三种形态与融合策略Mask 在实际工程里有几种不同的表达方式决定了 kernel 应该接收什么参数。第一种是加法 mask传进来的已经是最终的掩码值比如0.0表示保留-10000.0或-inf表示屏蔽。第二种是布尔 mask传进来的是True/Falsekernel 里根据布尔值决定要不要覆盖成-inf。第三种是根本不传 mask完全靠is_causal标志在 kernel 内部动态判断也就是根据当前处理的行号推导出 query 位置然后屏蔽掉未来位置的 key。我在算子设计里采用了“加法 mask 因果标志”的组合。理由很简单加法 mask 最通用你可以从任意布尔 mask 通过mask.to(x.dtype) * (-inf)转换得到而 causal 逻辑放在 kernel 内部可以省掉一份额外的 mask 张量和一次显存读写。如果mask传入的是空张量就认为无需 mask。2.3 一行一个block最直观且实用的线程组织方式接下来是 CUDA 编程里最核心的线程映射决策。Softmax 的归一化是逐行独立的每一行S个元素之间需要做一次全局归约。常见的方案有两种每个 block 处理一行block 内多个线程协作完成 reduce一个 warp 处理一行适合行宽较小的情况。对于 Attention 常见场景S通常是 128、256、512、1024甚至到 2048。一个 block 最多可以放 1024 个线程所以“每行一个 block”的方案在最常见范围内都能直接覆盖。如果S超过 1024也可以让每个线程按步长循环处理多个元素这样 block 数量仍然是行数线程数可以固定为 128 或 256。我采用的是固定blockDim128然后线程以tid为起始、以 128 为步长循环访问这一行内的所有元素。这样不管S是 128 还是 2048同一份代码都能跑只是每个线程处理的元素个数不同。行数rows B * H * S直接映射到gridDim.x每个 block 只需要知道自己是第几行然后从全局索引row * S开始处理逻辑非常清爽。3. 前向Kernel的实现一个可编译可运行的版本3.1 共享内存缓存与数据装载下面的前向 kernel 是我实际项目中采用的实现版本去掉了和具体业务耦合的部分保留了核心逻辑。为了方便讲解假设输入已经展开成二维视角rows B * H * Scols S。#include cuda_runtime.h #include math_constants.h template typename T __global__ void scaled_mask_softmax_forward_kernel( const T* __restrict__ x, const T* __restrict__ mask, T* __restrict__ y, const int rows, const int cols, const float scale, const bool has_mask, const bool is_causal) { const int row blockIdx.x; if (row rows) return; const int tid threadIdx.x; const int nthreads blockDim.x; const int q row % cols; extern __shared__ float sh[]; float* vals sh; float* red sh cols; float local_max -CUDART_INF_F; for (int i tid; i cols; i nthreads) { float v static_castfloat(x[row * cols i]) * scale; if (has_mask) { v static_castfloat(mask[row * cols i]); } if (is_causal i q) { v -CUDART_INF_F; } vals[i] v; local_max fmaxf(local_max, v); } red[tid] local_max; __syncthreads(); for (int s nthreads / 2; s 0; s 1) { if (tid s) { red[tid] fmaxf(red[tid], red[tid s]); } __syncthreads(); } const float row_max red[0]; __syncthreads(); float local_sum 0.0f; for (int i tid; i cols; i nthreads) { local_sum expf(vals[i] - row_max); } red[tid] local_sum; __syncthreads(); for (int s nthreads / 2; s 0; s 1) { if (tid s) { red[tid] red[tid s]; } __syncthreads(); } const float row_sum red[0]; for (int i tid; i cols; i nthreads) { y[row * cols i] static_castT(expf(vals[i] - row_max) / row_sum); } }这段代码里有几个细节值得单独拿出来说。第一我先把原始分数从全局显存读进共享内存vals后续的 max、sum、归一化都从共享内存取值而不是反复访问全局显存。共享内存的带宽远高于全局显存这是融合算子性能优势的主要来源之一。第二is_causal的判断放在 mask 之后。这是因为 causal 掩码本质上也是 mask如果之前叠加的 mask 已经屏蔽了某些位置那么继续覆盖成-inf不会改变语义。如果先判断 causal 再加 mask逻辑上也是等价的但要注意一旦i q后面的加法其实没有意义了所以放在后面更干净。第三extern __shared__ float sh[]是动态共享内存启动 kernel 时需要通过第三个配置参数指定大小。我采用的是(cols nthreads) * sizeof(float)前cols个 float 存放整行数据后nthreads个 float 作为归约缓冲区。3.2 两趟归约求max、求分母这个 kernel 的归约方式是最朴素的共享内存树形归约每次把线程数减半。red[tid]先保存每个线程的局部最大值然后第一轮tid 64的线程合并red[64]和red[0]第二轮tid 32合并red[32]和red[0]以此类推。整个过程需要log2(nthreads)次同步对 128 个线程来说就是 7 次代价可接受。这里有一个容易写错的地方在求完row_max之后紧接着把red缓冲区复用来求local_sum中间一定要加一次__syncthreads()。因为某个线程可能已经执行到red[tid] local_sum而另一个线程还没从red[0]里读出row_max这时候就会产生共享内存的读写竞争。我第一次写这个 kernel 的时候漏了这行同步结果在S512时偶尔出现nan排查了很久才意识到是同步问题。两趟归约在性能上并不是最优解因为要对整行数据做两遍遍历。更激进的方案是 online softmax也就是边读数据边维护 running max 和 running sum一遍遍历就能得到所有信息但代价是每个元素都要做一次乘法和除法来修正累积量。对于 seq_len 在 512 到 2048 这个区间的 Attention两趟遍历的共享内存访问开销其实很小代码反而更清晰易读。性能敏感到极致时再考虑换成 online 版本不迟。3.3 反向传播的雅可比推导与实现如果只做推理前向就够用了。但要接入训练必须实现反向 kernel。Softmax 的反向传播公式值得单独推导一遍因为它不是简单的“复制上游梯度”。记s_i x_i * scale mask_iy_i softmax(s_i)。对输出y的梯度是dy_i那么对s的梯度满足dx_i / ds_i 的形式 y_i * (dy_i - sum_j(dy_j * y_j))这个公式理解起来很直观Softmax 的雅可比不是对角矩阵因为输出之间互相影响归一化分母的存在意味着一个输出变大其他输出会相应变小。所以反向时要先计算一个全局的加权和dot sum_j(dy_j * y_j)然后每个位置减去这个公共项再乘上y_i。由于s_i x_i * scale mask_imask 是常数不产生梯度所以最终dx_i (dy_i - dot) * y_i * scale反向 kernel 就围绕这个公式展开。它只需要读前向保存的y和上游传来的dy先归约得到一个行内的dot再遍历一次写出dx。template typename T __global__ void scaled_mask_softmax_backward_kernel( const T* __restrict__ y, const T* __restrict__ dy, T* __restrict__ dx, const int rows, const int cols, const float scale) { const int row blockIdx.x; if (row rows) return; const int tid threadIdx.x; const int nthreads blockDim.x; extern __shared__ float sh[]; float* red sh; float local_dot 0.0f; for (int i tid; i cols; i nthreads) { float yv static_castfloat(y[row * cols i]); float dv static_castfloat(dy[row * cols i]); local_dot dv * yv; } red[tid] local_dot; __syncthreads(); for (int s nthreads / 2; s 0; s 1) { if (tid s) { red[tid] red[tid s]; } __syncthreads(); } const float dot red[0]; for (int i tid; i cols; i nthreads) { float yv static_castfloat(y[row * cols i]); float dv static_castfloat(dy[row * cols i]); dx[row * cols i] static_castT((dv - dot) * yv * scale); } }这里有个小技巧反向 kernel 不需要重新计算 softmax 的分母和 max因为y已经是归一化后的概率直接拿y参与梯度计算即可。这也是训练阶段必须在前向保存输出y的原因。有些实现会在反向里重新算一遍 softmax但那样纯粹是浪费显存带宽。4. 把它接进PyTorchC扩展与autograd链路4.1 从load_inline到setup.py两种工程化方式写好了 CUDA kernel接下来要让它能被 PyTorch 调用。PyTorch 提供了一套非常成熟的 C 扩展机制核心是torch.utils.cpp_extension。我建议在项目初期用load_inline做快速验证它不需要你维护复杂的setup.py直接传字符串源码即可。from torch.utils.cpp_extension import load_inline cpp_src ... C wrapper 代码 ... cuda_src ... CUDA kernel 代码 ... scaled_mask_softmax_ext load_inline( namescaled_mask_softmax_ext, cpp_sources[cpp_src], cuda_sources[cuda_src], functions[scaled_mask_softmax_forward, scaled_mask_softmax_backward], extra_cuda_cflags[-O3, --use_fast_math], verboseFalse, )当代码稳定下来需要纳入正式项目时再改成标准setup.py的方式。两种方式的切换成本很低核心的 C 和 CUDA 源码完全不用动。我个人的习惯是原型阶段load_inline一旦确认逻辑正确立刻转到setup.py因为后者对多文件组织、依赖声明、版本管理更友好。4.2 C Parser与类型分派C wrapper 是连接 PyTorch tensor 和 CUDA kernel 的桥梁。它要做的事情包括检查输入是否在 GPU 上、是否连续读取张量维度信息根据数据类型分派到对应的模板实例然后启动 kernel。#include torch/extension.h #include ATen/cuda/CUDAContext.h #define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x must be a CUDA tensor) #define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x must be contiguous) at::Tensor scaled_mask_softmax_forward( at::Tensor x, at::Tensor mask, double scale, bool is_causal) { CHECK_CUDA(x); CHECK_CONTIGUOUS(x); auto y at::empty_like(x); const int rows x.numel() / x.size(-1); const int cols x.size(-1); const bool has_mask mask.numel() 0; const scalar_t* mask_ptr has_mask ? mask.data_ptrscalar_t() : nullptr; AT_DISPATCH_FLOATING_TYPES_AND_HALF( x.scalar_type(), scaled_mask_softmax_forward, [] { auto stream at::cuda::getCurrentCUDAStream(); int threads 128; int smem (cols threads) * sizeof(float); scaled_mask_softmax_forward_kernelscalar_t rows, threads, smem, stream( x.data_ptrscalar_t(), mask_ptr, y.data_ptrscalar_t(), rows, cols, static_castfloat(scale), has_mask, is_causal); }); C10_CUDA_KERNEL_LAUNCH_CHECK(); return y; }有一个关键点x.numel() / x.size(-1)的计算方式使得这个算子天然支持[B, H, S, S]或[B*H, S, S]等不同维度的输入只要最后一维是 seq_len 就行。这给上层 Python 代码省去了很多 reshape 的操作。AT_DISPATCH_FLOATING_TYPES_AND_HALF是 PyTorch 提供的类型分派宏它会把torch.float32、torch.float64、torch.float16分别实例化对应的 kernel 模板。注意这里的scalar_t是宏展开时定义的局部类型名lambda 内部直接使用即可。这个宏还有一个好处是如果输入是torch.int64之类的类型编译时会直接抛错避免静默的类型错误。4.3 自定义Function与梯度接口Kernel 和 wrapper 都就绪之后最后一步是定义torch.autograd.Function。只有通过它PyTorch 才能在反向传播时自动调用我们的反向 kernel。import torch from torch.autograd import Function class ScaledMaskSoftmaxFunction(Function): staticmethod def forward(ctx, x, mask, scale, is_causal): x x.contiguous() if mask is None: mask torch.empty(0, devicex.device, dtypex.dtype) else: mask mask.contiguous() y scaled_mask_softmax_ext.scaled_mask_softmax_forward( x, mask, float(scale), bool(is_causal) ) ctx.scale float(scale) ctx.save_for_backward(y) return y staticmethod def backward(ctx, grad_output): (y,) ctx.saved_tensors grad_x scaled_mask_softmax_ext.scaled_mask_softmax_backward( y.contiguous(), grad_output.contiguous(), ctx.scale, ) return grad_x, None, None, None def scaled_mask_softmax(x, maskNone, scale1.0, is_causalFalse): return ScaledMaskSoftmaxFunction.apply(x, mask, scale, is_causal)这里有几个值得注意的细节。第一ctx.save_for_backward(y)保存的是前向输出而不是原始输入x。因为反向公式只需要y和上游梯度保存x反而是浪费显存。PyTorch 的save_for_backward机制会在反向结束后自动释放这些保存的张量不需要手动管理。第二backward里返回了四个值顺序和forward的输入参数一一对应分别是x、mask、scale、is_causal的梯度。由于mask、scale、is_causal都不需要梯度对应位置返回None。这个对应关系非常容易搞错如果你在backward里发现梯度形状对不上先检查返回值的数量和顺序。第三is_causal参数在forward里被保存到了ctx但这其实只是为了调试方便反向计算时并不需要它。因为 causal mask 的梯度本来就是零反向 kernel 只用y、dy、scale就够了。5. 验证正确性与测量性能先对齐语义再谈优化5.1 与标准实现逐项对拍写自定义算子最怕的就是“看起来对实际上错”。所以我强烈建议在开始性能测试之前先写一个严密的正确性验证脚本用 PyTorch 原生实现作为基准逐项对比。我的测试矩阵包括以下几个维度无 mask、无 causal纯带 scale 的 softmax带加法 maskmask 中同时包含0.0和-inf位置只带 causal 标志模拟 GPT 里的下三角掩码causal 和 mask 同时存在数据类型分别覆盖float32和float16seq_len 分别取 128、512、1024。torch.manual_seed(42) B, H, S, D 2, 4, 128, 64 x torch.randn(B, H, S, S, devicecuda) scale 1.0 / (D ** 0.5) mask torch.zeros(B, H, S, S, devicecuda) mask[:, :, :, S // 2:] -float(inf) y_ref torch.softmax(x * scale mask, dim-1) y_cus scaled_mask_softmax(x, maskmask, scalescale, is_causalFalse) print(max abs diff:, (y_ref - y_cus).abs().max().item())正常情况下max abs diff应该在1e-6量级。如果差异较大大概率是scale的传递类型出了问题或者 kernel 里的 mask 融合顺序不对。比如你传进来的 mask 是布尔型但 kernel 内部把它当加法 mask 直接相加True会被当成1.0结果自然不对。我前面提到过布尔 mask 必须先转换成0.0 / -inf的浮点形式再传入。5.2 CUDA事件计时别再用Python time性能测试不能用 Python 的time.time()因为它测到的是 CPU 侧的时间而 CUDA kernel 是异步执行的直接测 Python 时间会把 kernel 排队和同步的时间也算进去结果极不稳定。正确的做法是用torch.cuda.Event。start_event torch.cuda.Event(enable_timingTrue) end_event torch.cuda.Event(enable_timingTrue) # warm up for _ in range(10): y_cus scaled_mask_softmax(x, maskmask, scalescale) torch.cuda.synchronize() start_event.record() for _ in range(100): y_cus scaled_mask_softmax(x, maskmask, scalescale) end_event.record() torch.cuda.synchronize() print(average kernel time:, start_event.elapsed_time(end_event) / 100, ms)warm up 非常关键。GPU kernel 第一次执行时有初始化开销、缓存冷启动、cuDNN 或者 PyTorch 的 lazy initialization 等等直接把第一次调用计入统计会严重失真。我一般至少 warm up 10 次正式计时跑 100 次取平均。5.3 你可能看到“小样本反而更慢”的原因我很诚实地告诉你这个自定义算子在S128这种小尺寸上并不一定比 PyTorch 原生实现快。原因很现实kernel launch 本身有固定开销而且我们的 block 规模太小GPU 上大量计算单元处于闲置状态。真正能体现出融合优势的是S512甚至S1024的大规模场景。这时候全局内存访问次数的减少会显著拉低总耗时融合算子的收益才会变得肉眼可见。另外有一个 profiling 时的常见误区只看单个 kernel 的时间却不看端到端时间。PyTorch 原生实现是多个 kernel 串行它们之间的 launch 间隔和依赖等待同样耗时。自定义算子虽然单个 kernel 不一定是最快的但省掉了多 kernel 间的等待端到端往往有可观的收益。所以建议对比时既单独计时也对比一个完整的 Attention 前向加反向流程。6. 编译部署中我反复踩到的坑6.1 CUDA版本与PyTorch运行时不匹配自定义算子在本地跑通换了一台机器或者换了一个环境就编译失败这种问题我遇到过太多次。最典型的症状是编译时报nvcc版本和 PyTorch 编译时使用的 CUDA 版本不一致或者运行时报undefined symbol。PyTorch 的二进制发行版内部捆绑了一份 CUDA runtime它和系统里安装的 CUDA Toolkit 是两套东西。编译扩展时nvcc负责把 CUDA 源码编译成硬件代码而运行时的 CUDA runtime 库来自 PyTorch 内部。如果 PyTorch 是 CUDA 11.8 编译的系统里的nvcc却是 CUDA 12.3编译出来的cubin可能包含 PyTorch runtime 不认识的新特性跑起来就会报unknown error或者符号找不到。我的建议是在 conda 环境里用conda install cudatoolkit装和 PyTorch 匹配的 CUDA 版本然后通过CUDA_HOME指定对应的 toolkit 路径让nvcc跟 PyTorch 内部 runtime 保持同一个大版本。编译前先跑一句检查python -c import torch; print(torch.version.cuda) nvcc --version两个版本如果大版本不一致别急着调代码先把环境对齐再说。6.2 新显卡架构与TORCH_CUDA_ARCH_LIST如果你用的是比较新的显卡比如 RTX 40 系对应sm_90RTX 50 系对应sm_120而本机安装的 CUDA Toolkit 版本不够新编译时很可能会报ptxas fatal error: Unsupported gpu architecture。网上相关的报错信息也很多比如“sm_120 is not compatible”。这里的关键是理解TORCH_CUDA_ARCH_LIST环境变量。PyTorch 在编译扩展时会读取这个变量来决定为哪些 GPU 架构生成代码。如果没设置它会尝试自动探测当前显卡的 capability然后传给nvcc。问题在于老版本nvcc不认识新架构自动探测反而不安全。我的做法是显式指定一个兼容的架构列表比如export TORCH_CUDA_ARCH_LIST8.0;9.0这样nvcc就知道只需要生成 Ampere 和 Hopper 的代码不会试图去生成它根本不认识的sm_120。如果你希望代码在更多显卡上直接运行而不做 JIT可以把自己常用的架构都写进去但编译时间和二进制体积会相应增加。6.3 WSL2、多版本CUDA与nvcc搜索路径现在很多人在 Windows 上用 WSL2 搭 PyTorch 环境我也这么干过。WSL2 里最让人困惑的地方在于nvidia-smi显示的其实是 Windows 侧驱动的信息而nvcc --version显示的是 Linux 侧 CUDA Toolkit 的信息两者完全可以不同。驱动向上兼容只要驱动的版本不低于 toolkit 要求就行。多版本 CUDA 共存也是个高频话题。系统里装了两个 CUDA Toolkit 时/usr/local/cuda这个软链接指向哪个版本决定了默认nvcc是谁。我习惯这样做export CUDA_HOME/usr/local/cuda-12.4 export PATH$CUDA_HOME/bin:$PATH export LD_LIBRARY_PATH$CUDA_HOME/lib64:$LD_LIBRARY_PATH注意LD_LIBRARY_PATH不要和 conda 环境里的lib目录混在一起否则运行时可能加载到另一个版本的libcudart造成神秘的版本冲突。我遇到过最隐蔽的问题就是编译成功、加载成功但 kernel 启动之后计算结果完全错误最后发现是运行时加载了不同版本的 cudart 库。最后还有一个 fp16 相关的细节。AT_DISPATCH_FLOATING_TYPES_AND_HALF会把torch.float16也实例化出来但我们的 kernel 内部统一转成float做计算只在最终写回时转成half。这个设计是有意的fp16 的指数位太少如果中间累加和都用half存sum 和 exp 的误差会被放大得非常厉害。宁愿多花一点共享内存的转换开销也要保证数值精度。如果你在训练中发现 loss 和原版实现对比出现持续的小幅偏差可以试着把--use_fast_math去掉重新编译。这个编译选项会缩短expf等数学函数的精度大多数情况下没问题但在某些数据分布下会引入不可忽略的误差。从标题里的一个简单需求出发走到这里一个完整的自定义 ScaledMaskSoftmax 算子就已经接入训练流程了。回头看整个过程最花时间的其实不是写 kernel 本身而是搞清楚 PyTorch 的扩展机制、CUDA 的线程模型和共享内存的同步时序。每一步踩坑都对应着 CUDA 编程里最基础也最重要的概念搞清楚之后再去写别的融合算子比如 LayerNorm、GELU、Flash Attention 里的各种片段思路都是通用的。希望这篇文章能帮你少走几步弯路。
返回列表