news 2026/8/20 18:05:33

如何用 CuPy 编写自定义 CUDA 层?pytorch-pwc 相关层 4 个 Kernel 源码深度解读

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
如何用 CuPy 编写自定义 CUDA 层?pytorch-pwc 相关层 4 个 Kernel 源码深度解读

如何用 CuPy 编写自定义 CUDA 层?pytorch-pwc 相关层 4 个 Kernel 源码深度解读

【免费下载链接】pytorch-pwca reimplementation of PWC-Net in PyTorch that matches the official Caffe version项目地址: https://gitcode.com/gh_mirrors/py/pytorch-pwc

做光流估计、立体匹配的开发者常会遇到一个麻烦:PyTorch 官方并没有相关层(Correlation Layer),也就是论文中常说的 cost volume。pytorch-pwc 给出了一个非常优雅的答案——用 CuPy 编写自定义 CUDA 层,全程 Python,无需任何 C++ 编译配置。本文以 pytorch-pwc 的光流相关层为样本,深度解读其中 4 个 Kernel 的源码,帮你彻底搞懂如何用 CuPy 在 PyTorch 中写出高性能的自定义算子,看完就能照搬到自己的项目里。

一、pytorch-pwc 是什么?为什么需要自定义 CUDA 层?

pytorch-pwc 是 PWC-Net(CVPR 2018 论文《PWC-Net: CNNs for Optical Flow Using Pyramid, Warping, and Cost Volume》)的 PyTorch 重实现,核心目标是复刻官方 Caffe 版本的推理精度。整个网络遵循「金字塔特征提取 → 图像扭曲(Warping)→ 构建成本体(Cost Volume)→ 光流估计」的经典流程。

其中 cost volume 就是相关层的输出:对特征图的每个位置,在邻域 ±4 像素的 9×9 窗口内逐通道计算特征相关性,得到 81 个通道。这个算子太「专」了,PyTorch 没有现成实现,于是作者选择用 CuPy 手写 CUDA 内核。效果如何?官方 Caffe 与 PyTorch 实现的输出几乎完全一致:

![pytorch-pwc 光流结果对比:官方 Caffe 实现的光流图](https://raw.gitcode.com/gh_mirrors/py/pytorch-pwc/raw/fc2188815595b9fe3db94c7218c6f051eea0b012/comparison/official - caffe.png?utm_source=gitcode_repo_files)

![pytorch-pwc 光流结果对比:本仓库 PyTorch 实现的光流图](https://raw.gitcode.com/gh_mirrors/py/pytorch-pwc/raw/fc2188815595b9fe3db94c7218c6f051eea0b012/comparison/this - pytorch.png?utm_source=gitcode_repo_files)

二、为什么选 CuPy 而不是写 C++ 扩展?

写 C++/CUDA 扩展需要维护 setup.py、处理编译环境、为不同 CUDA 版本反复重编译,对新手非常不友好。而 CuPy 方案有三个肉眼可见的优势:

  • 纯 Python 编写,零编译配置:CUDA 内核以字符串形式直接写在 .py 文件里,安装好 cupy 即可运行。
  • RawKernel 直接执行 CUDA C:可以像写 .cu 文件一样编写内核,功能完整、性能不打折。
  • 自动缓存编译结果:配合cupy.memoize,同一内核只在首次调用时编译,之后直接复用,几乎零开销。

相关层完整代码都在 correlation/correlation.py,依赖也只需要在 requirements.txt 中加上 cupy 一行。

三、快速上手:安装与最小运行

克隆仓库后安装依赖即可运行(模型权重会自动下载):

git clone https://gitcode.com/gh_mirrors/py/pytorch-pwc cd pytorch-pwc pip install -r requirements.txt

用两张连续帧测试光流估计:

python run.py --model default --one ./images/one.png --two ./images/two.png --out ./out.flo

入口逻辑在 run.py,而核心的自定义 CUDA 层则被封装成correlation.FunctionCorrelation,在解码器中被反复调用(见 run.py)。

四、4 个 Kernel 源码深度解读

相关层一共包含4 个 CUDA Kernel:前向 2 个、反向 2 个。下面逐一拆解。

1️⃣ kernel_Correlation_rearrange:数据重排与补零

源码位置:correlation.py

这是前向的第一步:把标准 NCHW 布局的输入,重排成带 4 像素零填充的通道后置(HWC)布局,即[N, H+8, W+8, C]

extern "C" __global__ void kernel_Correlation_rearrange( const int n, const float* input, float* output) { int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; if (intIndex >= n) return; int intSample = blockIdx.z; // 样本索引 int intChannel = blockIdx.y; // 通道索引 ... // 每个像素平移到 (+4,+4) 的位置,四周自然补零 output[((intSample * (H+8) * (W+8) + intPaddedY * (W+8) + intPaddedX) * C) + intChannel] = fltValue; }

💡为什么补 4 像素零?因为后续相关计算要在 ±4 邻域内搜索,提前把边界外补成 0,之后所有内核都无需再做越界判断,大大简化代码。启动配置也很有意思:grid=(n/16, C, N),用blockIdx.y表示通道、blockIdx.z表示样本,三个维度各司其职。

2️⃣ kernel_Correlation_updateOutput:前向相关计算(核心)

源码位置:correlation.py

这是整个层的核心,负责计算 81 通道的 cost volume。启动配置为grid=(W, H, N)、block 32 个线程,并使用共享内存缓存 rbot0 的 patch,避免反复读显存:

// 把 rbot0 的 3D patch 载入共享内存 __shared__ float patch_data[...]; ... // 遍历 9x9 = 81 个位移 for (int top_channel = 0; top_channel < 81; top_channel++) { int s2o = top_channel % 9 - 4; // x 方向位移 int s2p = top_channel / 9 - 4; // y 方向位移 sum[ch_off] += patch_data[ch] * rbot1[idx2]; // 逐通道点积 ... top[index] = total_sum / (float)channels; // 按通道数归一化 }

💡性能技巧:共享内存的命中率直接决定该内核速度;32 个线程各自累加部分和,最后由ch_off == 0的线程做归约,把线程同步开销降到最低。

3️⃣ kernel_Correlation_updateGradOne:第一个输入的反向梯度

源码位置:correlation.py

反向传播时,需要分别计算两个输入的特征梯度。updateGradOne负责第一个输入(rbot0),关键设计有三点:

  • grid-stride 循环:用for (intIndex = blockIdx.x * blockDim.x + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x)让任意规模的张量都能被固定网格覆盖。
  • ROUND_OFF 取整技巧#define ROUND_OFF 50000,通过给负数加上大偏移再做整数除法,实现「负数的向上/向下取整」,从而精确算出每个像素能接收到梯度的 x/y 范围。
  • 边界裁剪xmin/xmax/ymin/ymax全部 clamp 到合法区间,越界部分直接跳过。

4️⃣ kernel_Correlation_updateGradTwo:第二个输入的反向梯度

源码位置:correlation.py

updateGradTwo的结构与updateGradOne几乎镜像:它读取的是 rbot0 而不是 rbot1,且位移方向相反(l - 4 - s2o),最终把梯度写回gradTwo。两者的循环与取整逻辑完全一致,可以对照阅读,非常利于理解「一个算子如何同时反向传播给两个输入」。

五、幕后魔法:SIZE_() 宏替换,一套代码适配所有尺寸

看源码时你可能会疑惑:内核里写的是SIZE_1(input)这种占位符,它怎么变成真实尺寸的?答案在 cupy_kernel 函数:

while True: objMatch = re.search('(SIZE_)([0-4])(\()([^\)]*)(\))', strKernel) if objMatch is None: break # 用张量的真实维度替换占位符 strKernel = strKernel.replace(objMatch.group(), str(intSizes[intArg]))

它用正则把SIZE_0()SIZE_4()替换为传入张量的实际尺寸,VALUE_则替换为 stride 公式。这样一来,同一份内核源码可以适配任意输入分辨率,不用为每个形状单独写代码,这也是本实现能在不同尺寸图像上稳定运行的关键。配合 cupy_launch 的cupy.memoize缓存,编译成本被压缩到一次。

六、4 个 Kernel 一览与复用建议

Kernel 名称阶段职责源码位置
kernel_Correlation_rearrange前向NCHW → 补零重排为 HWCcorrelation.py
kernel_Correlation_updateOutput前向计算 81 通道相关性correlation.py
kernel_Correlation_updateGradOne反向第一个输入的梯度correlation.py
kernel_Correlation_updateGradTwo反向第二个输入的梯度correlation.py

整个自定义 CUDA 层的封装思路,值得你在自己的项目里直接复用:

  1. 内核写成字符串,用SIZE_()占位符做形状无关化;
  2. cupy.RawKernel+cupy.memoize编译并缓存;
  3. 继承torch.autograd.Function,在forward中通过data_ptr()把 PyTorch 张量直接传给 CuPy 内核(见 forward 实现),在backward中调用两个梯度内核;
  4. 最后包一层torch.nn.Module对外暴露,使用体验与普通 PyTorch 层完全一致。

七、总结

通过 pytorch-pwc 的这 4 个 Kernel,你可以清楚看到用 CuPy 编写自定义 CUDA 层的完整套路:数据布局重排 → 共享内存加速 → 前后向成对实现 → 正则宏替换做形状无关化。这套模式不需要 C++ 工具链,却拥有接近手写 CUDA 的性能,非常适合在 PyTorch 中落地论文里的自定义算子。如果你正为「某个算子 PyTorch 没有」而发愁,不妨照着 pytorch-pwc 的样子,用 CuPy 写一个属于自己的自定义 CUDA 层。

【免费下载链接】pytorch-pwca reimplementation of PWC-Net in PyTorch that matches the official Caffe version项目地址: https://gitcode.com/gh_mirrors/py/pytorch-pwc

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/20 18:01:34

2026 西安 GEO 优化服务商口碑推荐:真实用户评价 + 核心优势 核验篇

核心结论 从合同边界角度看&#xff0c;判断服务商不能只看发稿数量&#xff0c;还要核验事实基础、方法完整性、数据口径和转化承接。更重要的是&#xff0c;当前西安企业常见的GEO路径主要包括内容布局、知识结构建设和AI品牌答案资产沉淀。广拓时代采用GAINS增长闭环组织GEO…

作者头像 李华
网站建设 2026/8/20 18:01:10

gruf 认证与安全:内置 Basic Auth 拦截器实战指南

gruf 认证与安全&#xff1a;内置 Basic Auth 拦截器实战指南 【免费下载链接】gruf gRPC Ruby Framework 项目地址: https://gitcode.com/gh_mirrors/gr/gruf 在微服务架构中&#xff0c;gRPC 服务的认证与安全是每个团队绕不开的话题。gruf 是一款优雅的 Ruby gRPC 框…

作者头像 李华
网站建设 2026/8/20 17:58:54

DeepSeek V4 Flash 实战指南:从 API 调用到本地部署的完整落地流程

1. 先搞清楚 DeepSeek V4 Flash 到底是个什么定位最近关于 DeepSeek V4 Flash 的讨论很多&#xff0c;标题里“杀疯了”、“跻身前五”这些说法&#xff0c;核心指向一个事实&#xff1a;这是一个在性能、成本和速度之间找到了新平衡点的模型。它不是那个参数规模最大、能力最强…

作者头像 李华
网站建设 2026/8/20 17:58:25

FF14终极启动器XIVLauncher完全指南

FF14终极启动器XIVLauncher完全指南 【免费下载链接】FFXIVQuickLauncher Custom launcher for FFXIV 项目地址: https://gitcode.com/GitHub_Trending/ff/FFXIVQuickLauncher 凌晨一点&#xff0c;进度条停在98%已经快半小时——这是今晚第三次了。XIVLauncher&#xf…

作者头像 李华