阶段 5 · 工程化与 AI14 / 15约 26 分钟
接入 PyTorch:把你的 kernel 变成一个算子
从 load_inline 快速验证到正式的 setuptools 扩展与自动求导
学完这一课你会
- 用 load_inline 在几分钟内验证一个自定义 kernel
- 写出规范的 setuptools 扩展并处理好张量检查
- 为自定义算子接上 autograd
- 了解 torch.compile 时代自定义算子的正确注册方式
学 CUDA 的一个非常现实的落脚点,是给 PyTorch 写自定义算子。PyTorch 的 C++ 扩展机制让这件事比想象中简单:你的 kernel 只需要接收 `torch::Tensor`,其余的内存管理、设备调度、类型分发都由框架处理。
最快路径:load_inline
Python
1import torch2from torch.utils.cpp_extension import load_inline3 4cuda_src = r"""5#include <torch/extension.h>6 7__global__ void square_kernel(const float* in, float* out, int n) {8 int i = blockIdx.x * blockDim.x + threadIdx.x;9 if (i < n) out[i] = in[i] * in[i];10}11 12torch::Tensor square(torch::Tensor x) {13 TORCH_CHECK(x.is_cuda(), "输入必须在 CUDA 上");14 TORCH_CHECK(x.is_contiguous(), "输入必须连续");15 TORCH_CHECK(x.scalar_type() == torch::kFloat32, "目前只支持 float32");16 17 auto out = torch::empty_like(x);18 int n = x.numel();19 int threads = 256, blocks = (n + threads - 1) / threads;20 21 square_kernel<<<blocks, threads>>>(22 x.data_ptr<float>(), out.data_ptr<float>(), n);23 24 // 用 PyTorch 当前流,保证和框架其他操作的顺序正确25 C10_CUDA_KERNEL_LAUNCH_CHECK();26 return out;27}28"""29 30cpp_src = "torch::Tensor square(torch::Tensor x);"31 32mod = load_inline(33 name="my_square",34 cpp_sources=cpp_src,35 cuda_sources=cuda_src,36 functions=["square"],37 extra_cuda_cflags=["-O3", "--use_fast_math"],38 verbose=True,39)40 41x = torch.randn(1_000_000, device="cuda")42torch.testing.assert_close(mod.square(x), x * x)43print("通过")正式做法:setuptools 扩展
C++
1#include <torch/extension.h>2 3// CUDA 侧的实现声明(定义在 .cu 文件里)4torch::Tensor rmsnorm_cuda(torch::Tensor x, torch::Tensor weight, double eps);5 6#define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " 必须是 CUDA 张量")7#define CHECK_CONTIG(x) TORCH_CHECK(x.is_contiguous(), #x " 必须连续")8#define CHECK_INPUT(x) CHECK_CUDA(x); CHECK_CONTIG(x)9 10torch::Tensor rmsnorm(torch::Tensor x, torch::Tensor weight, double eps) {11 CHECK_INPUT(x);12 CHECK_INPUT(weight);13 TORCH_CHECK(x.size(-1) == weight.size(0), "最后一维必须与 weight 长度一致");14 return rmsnorm_cuda(x, weight, eps);15}16 17PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {18 m.def("rmsnorm", &rmsnorm, "RMSNorm (CUDA)",19 py::arg("x"), py::arg("weight"), py::arg("eps") = 1e-6);20}CUDA C++
1#include <torch/extension.h>2#include <c10/cuda/CUDAStream.h>3 4template <typename scalar_t>5__global__ void rmsnorm_kernel(const scalar_t* __restrict__ x,6 const scalar_t* __restrict__ w,7 scalar_t* __restrict__ out,8 int hidden, float eps) {9 int row = blockIdx.x; // 一个 block 负责一行10 const scalar_t* xr = x + (size_t)row * hidden;11 scalar_t* orow = out + (size_t)row * hidden;12 13 // 第一遍:算平方和14 float sum = 0.0f;15 for (int i = threadIdx.x; i < hidden; i += blockDim.x) {16 float v = static_cast<float>(xr[i]);17 sum += v * v;18 }19 // warp 内归约 → 跨 warp 归约20 for (int off = 16; off > 0; off >>= 1)21 sum += __shfl_down_sync(0xffffffff, sum, off);22 23 __shared__ float warpSums[32];24 int lane = threadIdx.x & 31, wid = threadIdx.x >> 5;25 if (lane == 0) warpSums[wid] = sum;26 __syncthreads();27 28 if (wid == 0) {29 sum = (lane < (blockDim.x + 31) / 32) ? warpSums[lane] : 0.0f;30 for (int off = 16; off > 0; off >>= 1)31 sum += __shfl_down_sync(0xffffffff, sum, off);32 if (lane == 0) warpSums[0] = rsqrtf(sum / hidden + eps);33 }34 __syncthreads();35 float scale = warpSums[0];36 37 // 第二遍:归一化并乘上权重38 for (int i = threadIdx.x; i < hidden; i += blockDim.x) {39 orow[i] = static_cast<scalar_t>(static_cast<float>(xr[i]) * scale40 * static_cast<float>(w[i]));41 }42}43 44torch::Tensor rmsnorm_cuda(torch::Tensor x, torch::Tensor weight, double eps) {45 auto out = torch::empty_like(x);46 int hidden = x.size(-1);47 int rows = x.numel() / hidden;48 49 int threads = std::min(1024, ((hidden + 31) / 32) * 32);50 auto stream = at::cuda::getCurrentCUDAStream(); // 用 PyTorch 的流!51 52 AT_DISPATCH_FLOATING_TYPES_AND2(53 at::ScalarType::Half, at::ScalarType::BFloat16,54 x.scalar_type(), "rmsnorm_cuda", [&] {55 rmsnorm_kernel<scalar_t><<<rows, threads, 0, stream>>>(56 x.data_ptr<scalar_t>(), weight.data_ptr<scalar_t>(),57 out.data_ptr<scalar_t>(), hidden, static_cast<float>(eps));58 });59 return out;60}Python
1from setuptools import setup2from torch.utils.cpp_extension import BuildExtension, CUDAExtension3 4setup(5 name="my_ops",6 ext_modules=[7 CUDAExtension(8 name="my_ops._C",9 sources=["src/ops.cpp", "src/rmsnorm.cu"],10 extra_compile_args={11 "cxx": ["-O3"],12 "nvcc": ["-O3", "--use_fast_math",13 "-gencode", "arch=compute_80,code=sm_80",14 "-gencode", "arch=compute_86,code=sm_86"],15 },16 )17 ],18 cmdclass={"build_ext": BuildExtension},19)接上 autograd
Python
1import torch2from my_ops import _C3 4class RMSNormFn(torch.autograd.Function):5 @staticmethod6 def forward(ctx, x, weight, eps=1e-6):7 out = _C.rmsnorm(x, weight, eps)8 ctx.save_for_backward(x, weight)9 ctx.eps = eps10 return out11 12 @staticmethod13 def backward(ctx, grad_out):14 x, weight = ctx.saved_tensors15 # 生产环境这里应该也是一个 CUDA kernel;16 # 开发阶段先用 PyTorch 算子实现,确保数值正确后再替换。17 grad_x, grad_w = _C.rmsnorm_backward(grad_out, x, weight, ctx.eps)18 return grad_x, grad_w, None19 20def rmsnorm(x, weight, eps=1e-6):21 return RMSNormFn.apply(x, weight, eps)22 23# 数值梯度校验:新算子必做24x = torch.randn(4, 128, device="cuda", dtype=torch.double, requires_grad=True)25w = torch.randn(128, device="cuda", dtype=torch.double, requires_grad=True)26assert torch.autograd.gradcheck(rmsnorm, (x, w), eps=1e-6, atol=1e-4)torch.compile 时代:注册为自定义算子
如果你的模型要走 torch.compile,直接调用 pybind 出来的函数会造成图断裂(graph break),编译器无法把它纳入优化。正确做法是用 torch.library 把算子注册进 PyTorch 的算子体系,并提供一个 FakeTensor 实现(描述输出的形状和 dtype,供图捕获阶段推导)。
Python
1import torch2from torch.library import custom_op, register_fake3 4@custom_op("my_ops::rmsnorm", mutates_args=())5def rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:6 from my_ops import _C7 return _C.rmsnorm(x, weight, eps)8 9@register_fake("my_ops::rmsnorm")10def _(x, weight, eps=1e-6):11 # 只需描述输出的元信息,不做真实计算12 return torch.empty_like(x)13 14# 现在它能被 torch.compile 正常捕获,不会 graph break15compiled = torch.compile(lambda a, b: torch.ops.my_ops.rmsnorm(a, b))自测
先自己在心里回答一遍,再展开对照。答不上来的说明这一段值得重读。