拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
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],当S=1024、B=8、H=32时,光这一个张量就是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_i=100,单精度浮点数直接溢出成inf。Attention 分数经过 scale 之后通常在一个可控范围内,但加了 mask 或者训练初期参数不稳定时,出现极端值是常有的事。所以数值稳定版的 softmax 一定会做一步变换:

m = max_j(score_j) output_i = exp(score_i - m) / sum_j(exp(score_j - m))

减掉行最大值之后,所有指数项的输入都不超过 0,exp的结果落在(0, 1]区间,彻底避免溢出。这个虽然是很基础的常识,但写 CUDA kernel 的时候特别容易漏:一旦漏掉,可能在特定输入下产生nan,而且这种 bug 非常难查。我的建议是,不管你认为输入范围多安全,稳定版变换一定要做。

2.2 mask的三种形态与融合策略

Mask 在实际工程里有几种不同的表达方式,决定了 kernel 应该接收什么参数。

第一种是加法 mask,传进来的已经是最终的掩码值,比如0.0表示保留,-10000.0或-inf表示屏蔽。第二种是布尔 mask,传进来的是True/False,kernel 里根据布尔值决定要不要覆盖成-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。

我采用的是固定blockDim=128,然后线程以tid为起始、以 128 为步长循环访问这一行内的所有元素。这样不管S是 128 还是 2048,同一份代码都能跑,只是每个线程处理的元素个数不同。行数rows = B * H * S,直接映射到gridDim.x,每个 block 只需要知道自己是第几行,然后从全局索引row * S开始处理,逻辑非常清爽。

3. 前向Kernel的实现:一个可编译可运行的版本

3.1 共享内存缓存与数据装载

下面的前向 kernel 是我实际项目中采用的实现版本,去掉了和具体业务耦合的部分,保留了核心逻辑。为了方便讲解,假设输入已经展开成二维视角:rows = B * H * S,cols = 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_cast<float>(x[row * cols + i]) * scale; if (has_mask) { v += static_cast<float>(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_cast<T>(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 的时候漏了这行同步,结果在S=512时偶尔出现nan,排查了很久才意识到是同步问题。

两趟归约在性能上并不是最优解,因为要对整行数据做两遍遍历。更激进的方案是 online softmax,也就是边读数据边维护 running max 和 running sum,一遍遍历就能得到所有信息,但代价是每个元素都要做一次乘法和除法来修正累积量。对于 seq_len 在 512 到 2048 这个区间的 Attention,两趟遍历的共享内存访问开销其实很小,代码反而更清晰易读。性能敏感到极致时,再考虑换成 online 版本不迟。

3.3 反向传播的雅可比推导与实现

如果只做推理,前向就够用了。但要接入训练,必须实现反向 kernel。Softmax 的反向传播公式值得单独推导一遍,因为它不是简单的“复制上游梯度”。

记s_i = x_i * scale + mask_i,y_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_i,mask 是常数不产生梯度,所以最终:

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_cast<float>(y[row * cols + i]); float dv = static_cast<float>(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_cast<float>(y[row * cols + i]); float dv = static_cast<float>(dy[row * cols + i]); dx[row * cols + i] = static_cast<T>((dv - dot) * yv * scale); } }

这里有个小技巧:反向 kernel 不需要重新计算 softmax 的分母和 max,因为y已经是归一化后的概率,直接拿y参与梯度计算即可。这也是训练阶段必须在前向保存输出y的原因。有些实现会在反向里重新算一遍 softmax,但那样纯粹是浪费显存带宽。

4. 把它接进PyTorch:C++扩展与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( name="scaled_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"], verbose=False, )

当代码稳定下来,需要纳入正式项目时,再改成标准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_ptr<scalar_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_kernel<scalar_t> <<<rows, threads, smem, stream>>>( x.data_ptr<scalar_t>(), mask_ptr, y.data_ptr<scalar_t>(), rows, cols, static_cast<float>(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, device=x.device, dtype=x.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, mask=None, scale=1.0, is_causal=False): 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;
  • 带加法 mask,mask 中同时包含0.0和-inf位置;
  • 只带 causal 标志,模拟 GPT 里的下三角掩码;
  • causal 和 mask 同时存在;
  • 数据类型分别覆盖float32和float16;
  • seq_len 分别取 128、512、1024。
torch.manual_seed(42) B, H, S, D = 2, 4, 128, 64 x = torch.randn(B, H, S, S, device="cuda") scale = 1.0 / (D ** 0.5) mask = torch.zeros(B, H, S, S, device="cuda") mask[:, :, :, S // 2:] = -float("inf") y_ref = torch.softmax(x * scale + mask, dim=-1) y_cus = scaled_mask_softmax(x, mask=mask, scale=scale, is_causal=False) 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_timing=True) end_event = torch.cuda.Event(enable_timing=True) # warm up for _ in range(10): y_cus = scaled_mask_softmax(x, mask=mask, scale=scale) torch.cuda.synchronize() start_event.record() for _ in range(100): y_cus = scaled_mask_softmax(x, mask=mask, scale=scale) 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 你可能看到“小样本反而更慢”的原因

我很诚实地告诉你,这个自定义算子在S=128这种小尺寸上,并不一定比 PyTorch 原生实现快。原因很现实:kernel launch 本身有固定开销,而且我们的 block 规模太小,GPU 上大量计算单元处于闲置状态。真正能体现出融合优势的,是S>=512甚至S>=1024的大规模场景。这时候全局内存访问次数的减少会显著拉低总耗时,融合算子的收益才会变得肉眼可见。

另外有一个 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_90,RTX 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_LIST="8.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 里的各种片段,思路都是通用的。希望这篇文章能帮你少走几步弯路。

返回列表