如何用 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-NetCVPR 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_sourcegitcode_repo_files)![pytorch-pwc 光流结果对比本仓库 PyTorch 实现的光流图](https://raw.gitcode.com/gh_mirrors/py/pytorch-pwc/raw/fc2188815595b9fe3db94c7218c6f051eea0b012/comparison/this - pytorch.png?utm_sourcegitcode_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, H8, W8, 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 * (H8) * (W8) intPaddedY * (W8) 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.pyupdateGradTwo的结构与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.pykernel_Correlation_updateOutput前向计算 81 通道相关性correlation.pykernel_Correlation_updateGradOne反向第一个输入的梯度correlation.pykernel_Correlation_updateGradTwo反向第二个输入的梯度correlation.py整个自定义 CUDA 层的封装思路值得你在自己的项目里直接复用内核写成字符串用SIZE_()占位符做形状无关化用cupy.RawKernelcupy.memoize编译并缓存继承torch.autograd.Function在forward中通过data_ptr()把 PyTorch 张量直接传给 CuPy 内核见 forward 实现在backward中调用两个梯度内核最后包一层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),仅供参考

相关新闻

基于SpringBoot+Vue的美妆产品推荐系统(源码+讲解视频+LW)

基于SpringBoot+Vue的美妆产品推荐系统(源码+讲解视频+LW)

联系博主 温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台官方提供的学长联系方式的名片! 温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台官方提供的学长联系方式的名片! 温馨提示:本人主页置顶文章(点我)开头有 …

2026/8/20 18:05:42 阅读更多 →
uivonim源码剖析:从instance-api到window-manager,一份Neovim GUI架构阅读指南

uivonim源码剖析:从instance-api到window-manager,一份Neovim GUI架构阅读指南

uivonim源码剖析:从instance-api到window-manager,一份Neovim GUI架构阅读指南 【免费下载链接】uivonim Fork of the Veonim Neovim GUI 项目地址: https://gitcode.com/gh_mirrors/ui/uivonim uivonim 是 Veonim Neovim GUI 的活跃 fork&#x…

2026/8/21 18:47:37 阅读更多 →
暗黑2存档编辑器 d2s-editor 上手攻略:免费开源,拖入 .d2s 就能改出满级角色

暗黑2存档编辑器 d2s-editor 上手攻略:免费开源,拖入 .d2s 就能改出满级角色

暗黑2存档编辑器 d2s-editor 上手攻略:免费开源,拖入 .d2s 就能改出满级角色 【免费下载链接】d2s-editor 项目地址: https://gitcode.com/gh_mirrors/d2/d2s-editor 你有没有过这样的时刻:想重温《暗黑破坏神2》,却实在不…

2026/8/21 18:57:34 阅读更多 →

最新新闻

Obsidian插件LyricFlux:在笔记中无缝集成音乐管理与播放

Obsidian插件LyricFlux:在笔记中无缝集成音乐管理与播放

你是否曾有过这样的体验:在 Obsidian 中整理读书笔记或写作时,想听一首歌来激发灵感,却不得不切换到音乐播放器,打断心流?或者,你收藏了大量音乐相关的资料、歌词、乐评,却无法将它们与你本地的…

2026/8/21 20:10:09 阅读更多 →
三步跑起 BETAFPV Configurator 地面站:克隆到绑定全打通

三步跑起 BETAFPV Configurator 地面站:克隆到绑定全打通

三步跑起 BETAFPV Configurator 地面站:克隆到绑定全打通 【免费下载链接】BETAFPV_Configurator 项目地址: https://gitcode.com/gh_mirrors/be/BETAFPV_Configurator BETAFPV Configurator 是一套开源地面站工具,负责 BETAFPV 遥控器与飞控的固…

2026/8/21 20:10:09 阅读更多 →
开源Web打印设计器OpenPrint:可视化拖拽与数据绑定实战

开源Web打印设计器OpenPrint:可视化拖拽与数据绑定实战

这次我们来看一个开源的 Web 打印报表设计器——OpenPrint。如果你正在寻找一个能嵌入到业务系统里、支持可视化拖拽设计、并能将数据绑定到模板上生成 PDF 或直接打印的解决方案,那这个项目值得你花十分钟了解一下。 OpenPrint 的核心是让 Web 端的报表设计和打印…

2026/8/21 20:10:09 阅读更多 →
留痕保姆级指南:微信聊天记录导出与个人数据管理,免费本地提取让隐私自己做主

留痕保姆级指南:微信聊天记录导出与个人数据管理,免费本地提取让隐私自己做主

留痕保姆级指南:微信聊天记录导出与个人数据管理,免费本地提取让隐私自己做主 【免费下载链接】WeChatMsg 提取微信聊天记录,将其导出成HTML、Word、CSV文档永久保存,对聊天记录进行分析生成年度聊天报告 项目地址: https://git…

2026/8/21 20:10:09 阅读更多 →
3步跑通 League Akari:英雄联盟智能选与对局自动化客户端的完整上手指南

3步跑通 League Akari:英雄联盟智能选与对局自动化客户端的完整上手指南

3步跑通 League Akari:英雄联盟智能选与对局自动化客户端的完整上手指南 【免费下载链接】League-Toolkit An all-in-one toolkit for LeagueClient. Gathering power 🚀. 项目地址: https://gitcode.com/gh_mirrors/le/League-Toolkit 你是否在选…

2026/8/21 20:10:09 阅读更多 →
兄弟打印机E3错误深度解析:打印正常却报错的传感器故障排查指南

兄弟打印机E3错误深度解析:打印正常却报错的传感器故障排查指南

兄弟T725DW打印机,打印效果一切正常,纸张顺畅进出,墨迹清晰,但控制面板或电脑端却固执地弹出一个“无法打印E3”的错误。你反复检查了纸张、墨盒、连接线,甚至重启了无数次,问题依旧。这不是一个功能性的“…

2026/8/21 20:09:08 阅读更多 →

日新闻

机场边检旅客定位系统国产化白皮书:算法、硬件、底座平台全程自主

机场边检旅客定位系统国产化白皮书:算法、硬件、底座平台全程自主

前言随着国家数字基础设施信创替代、关键技术自主可控战略持续深化,口岸智慧安防、边检智能管控领域正全面进入国产化、自主化、安全可控升级周期。当前国内机场边检旅客识别与定位体系长期依赖国外商用视觉算法、进口成像硬件、闭源通用计算平台,存在核…

2026/8/21 0:00:42 阅读更多 →
别再把“数字孪生”当空间智能了!镜像视界揭开四维时空的真正面纱

别再把“数字孪生”当空间智能了!镜像视界揭开四维时空的真正面纱

别再把“数字孪生”当空间智能了!镜像视界揭开四维时空的真正面纱当下数字化建设浪潮中,很多项目将三维可视化、视频贴图叠加的数字孪生等同于空间智能。传统数字孪生更多停留在三维场景复刻,擅长把物理世界“画出来、展示出来”,…

2026/8/21 0:00:42 阅读更多 →
105、车载温度范围-40°C到85°C的影像质量一致性——ISP参数温漂补偿与产线标定策略

105、车载温度范围-40°C到85°C的影像质量一致性——ISP参数温漂补偿与产线标定策略

105、车载温度范围-40C到85C的影像质量一致性——ISP参数温漂补偿与产线标定策略 去年冬天在北方某车厂做A样评审,凌晨四点的黑河试验场,零下三十三度。客户拿了一台冷启动的车,中控屏上倒车影像全是雪花噪点,暗部细节直接糊成一片。我第一反应是sensor温度没上来,暗电流…

2026/8/21 0:00:42 阅读更多 →

周新闻

基于阿里云与通义千问(Qwen)构建AI应用:从模型调用到生产部署的完整实践指南

基于阿里云与通义千问(Qwen)构建AI应用:从模型调用到生产部署的完整实践指南

如果你是一名开发者,最近可能已经感受到了AI大模型正在从“玩具”变成“生产力工具”的强烈信号。从代码补全到智能Agent,从本地部署到云端API,我们正处在一个技术栈快速重构的节点。然而,面对层出不穷的模型、框架和工具&#xf…

2026/8/21 3:21:33 阅读更多 →
工业通信系统底层逻辑:04 反射——高频能量撞墙之后会发生什么?

工业通信系统底层逻辑:04 反射——高频能量撞墙之后会发生什么?

第四篇:反射——高频能量撞墙之后会发生什么? —— 你以为信号已经过去了,其实它正在回来打你 老Q的现场笔记 第五季,我们正式进入工业神经系统层。这里不再是单个设备的战斗,而是整个工厂“经脉”层面的秩序之战。从这一篇开始,你将第一次看清:看似简单的信号传播,背…

2026/8/21 0:02:09 阅读更多 →
【文章复现】非线性值迭代自适应动态规划(ADP):离散时间非线性系统的策略迭代自适应动态规划算法研究附Matlab代码

【文章复现】非线性值迭代自适应动态规划(ADP):离散时间非线性系统的策略迭代自适应动态规划算法研究附Matlab代码

✅作者简介:热爱科研的Matlab仿真开发者,擅长毕业设计辅导、数学建模、数据处理、建模仿真、程序设计、完整代码获取、论文复现及科研仿真。🍎 往期回顾关注个人主页:Matlab科研工作室👇 关注我领取海量matlab电子书和…

2026/8/21 6:07:56 阅读更多 →

月新闻

免费解锁百度网盘SVIP加速:macOS用户必备的下载提速终极指南

免费解锁百度网盘SVIP加速:macOS用户必备的下载提速终极指南

免费解锁百度网盘SVIP加速:macOS用户必备的下载提速终极指南 【免费下载链接】BaiduNetdiskPlugin-macOS For macOS.百度网盘 破解SVIP、下载速度限制~ 项目地址: https://gitcode.com/gh_mirrors/ba/BaiduNetdiskPlugin-macOS 还在为百度网盘macOS版的龟速下…

2026/8/21 16:42:28 阅读更多 →
终极ncmdump指南:3分钟实现网易云NCM音乐解密与格式转换

终极ncmdump指南:3分钟实现网易云NCM音乐解密与格式转换

终极ncmdump指南:3分钟实现网易云NCM音乐解密与格式转换 【免费下载链接】ncmdump 项目地址: https://gitcode.com/gh_mirrors/ncmd/ncmdump 还在为网易云音乐下载的NCM格式文件无法在其他播放器播放而烦恼吗?ncmdump解密工具帮你轻松解决这个困…

2026/8/20 21:46:49 阅读更多 →
HarmonyOS 应用开发《掌上英语》第81篇: 智能体卡片:为英语学习 App 打造桌面级学习助手

HarmonyOS 应用开发《掌上英语》第81篇: 智能体卡片:为英语学习 App 打造桌面级学习助手

AgentCard 智能体卡片:为英语学习 App 打造桌面级学习助手适用平台:HarmonyOS 7.0 (API 26 Beta)一、引言 HarmonyOS 7.0(API 26 Beta)新增了 AgentCard 智能体卡片能力,这是继 HMAF(鸿蒙智能体框架&#x…

2026/8/21 0:14:22 阅读更多 →