<<<CUDA C++ 学习路线grid · block · warp · lane
阶段 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 扩展

src/ops.cpp — 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}
src/rmsnorm.cu — 一行一个 block 的实现
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}
setup.py
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

自定义 Function
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,供图捕获阶段推导)。

torch.library 注册
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))

自测

先自己在心里回答一遍,再展开对照。答不上来的说明这一段值得重读。