ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

算子融合(Kernel Fusion):用做菜的比喻彻底讲明白

算子融合(Kernel Fusion):用做菜的比喻彻底讲明白 算子融合Kernel Fusion用做菜的比喻彻底讲明白引言一个让你秒懂的比喻想象你在厨房做一道菜——炒鸡蛋番茄。笨办法未融合1. 打鸡蛋 → 装进碗里 → 放进冰箱 2. 从冰箱拿出鸡蛋 → 炒熟 → 装进盘子 → 放进冰箱 3. 从冰箱拿出炒蛋 → 加番茄 → 装盘 → 上桌看出问题了吗你在反复地放冰箱、拿冰箱每一步都要跑到冰箱那里存取食材浪费大量时间。聪明办法融合打鸡蛋 → 炒熟 → 加番茄 → 直接上桌一气呵成不需要中途放冰箱这就是算子融合的核心思想。GPU 里的冰箱是什么在 GPU 编程中做菜比喻GPU 对应速度冰箱远存取慢全局内存显存慢 手边的操作台近用起来快寄存器快 每次跑去冰箱访问全局内存耗时关键事实从显存读写数据非常慢相对而言在寄存器里计算非常快性能瓶颈通常不是算得慢而是存取数据慢用代码看清楚问题场景先平方再开方我们要对一组数据做两个操作先算x²再算√结果。❌ 笨办法未融合thrust::device_vectorfloatdata(1000);// 操作 1平方thrust::transform(data.begin(),data.end(),data.begin(),[]__device__(floatx){returnx*x;});// 操作 2开方thrust::transform(data.begin(),data.end(),data.begin(),[]__device__(floatx){returnsqrtf(x);});GPU 实际在做什么第一个 transform平方 ┌─────────────────────────────┐ │ 1. 从显存读取 x 慢 │ │ 2. 计算 x * x 快 │ │ 3. 写回显存 慢 │ └─────────────────────────────┘ 第二个 transform开方 ┌─────────────────────────────┐ │ 4. 从显存读取结果 慢 │ ← 又跑去冰箱 │ 5. 计算 sqrt 快 │ │ 6. 写回显存 慢 │ └─────────────────────────────┘ 总共4 次显存访问2 读 2 写✅ 聪明办法融合thrust::device_vectorfloatdata(1000);// 一次搞定两个操作thrust::transform(data.begin(),data.end(),data.begin(),[]__device__(floatx){floattmpx*x;// 平方在寄存器里returnsqrtf(tmp);// 开方在寄存器里});GPU 实际在做什么一个 transform平方开方 ┌─────────────────────────────┐ │ 1. 从显存读取 x 慢 │ │ 2. 计算 x * x 快 │ ← 结果留在寄存器 │ 3. 计算 sqrt 快 │ ← 直接用寄存器的值 │ 4. 写回显存 慢 │ └─────────────────────────────┘ 总共2 次显存访问1 读 1 写对比未融合4 次显存访问融合后2 次显存访问省了一半的跑冰箱时间为什么中间结果不需要存显存这是理解融合的核心关键未融合中间结果被迫存显存// 第一步算完结果必须写回显存thrust::transform(...,square);// x² 存到显存// 第二步才能从显存读取thrust::transform(...,sqrt);// 从显存读 x²再开方为什么因为这是两个独立的函数调用就像两次独立的做菜过程。第一次做完必须把成品放好存显存第二次才能拿出来继续加工。融合中间结果留在寄存器thrust::transform(...,[](floatx){floattmpx*x;// tmp 是临时变量存在寄存器里returnsqrtf(tmp);// 直接用寄存器里的 tmp});为什么快因为tmp只是一个临时变量它就在 GPU 核心的寄存器里用完即走根本不需要跑去显存一张图看懂融合前后未融合数据反复进出显存显存 GPU核心寄存器 ┌────┐ ┌──────────┐ │ x │ ──读取──────────→ │ x │ │ │ │ ↓ │ │ │ │ x*x │ │x² │ ←──写回─────────── │ │ ├────┤ ├──────────┤ │x² │ ──读取──────────→ │ x² │ │ │ │ ↓ │ │ │ │ sqrt(x²) │ │√x² │ ←──写回─────────── │ │ └────┘ └──────────┘ ↑ 4次进出显存融合后数据只进出一次显存 GPU核心寄存器 ┌────┐ ┌──────────┐ │ x │ ──读取──────────→ │ x │ │ │ │ ↓ │ │ │ │ x*x │ ← 留在寄存器 │ │ │ ↓ │ │ │ │ sqrt(..) │ ← 继续用寄存器 │√x² │ ←──写回─────────── │ │ └────┘ └──────────┘ ↑ 只2次进出显存生活中更多的融合例子例子 1数据加工流水线任务把每个数字1再*2再-3❌ 笨办法thrust::transform(data.begin(),data.end(),data.begin(),[]__device__(floatx){returnx1;});// 存显存thrust::transform(data.begin(),data.end(),data.begin(),[]__device__(floatx){returnx*2;});// 存显存thrust::transform(data.begin(),data.end(),data.begin(),[]__device__(floatx){returnx-3;});// 存显存// 6 次显存访问3 读 3 写✅ 聪明办法thrust::transform(data.begin(),data.end(),data.begin(),[]__device__(floatx){return(x1)*2-3;// 一次搞定});// 2 次显存访问1 读 1 写省了 3 倍的显存访问例子 2求和前先平方transform_reduce任务计算所有数字的平方和Σx²❌ 笨办法// 步骤 1先平方存到临时数组thrust::device_vectorfloatsquared(N);thrust::transform(data.begin(),data.end(),squared.begin(),[]__device__(floatx){returnx*x;});// 读写// 步骤 2再求和floatsumthrust::reduce(squared.begin(),squared.end());// 读// 3 次显存访问还浪费了一个临时数组的内存✅ 聪明办法用融合专用算子// transform_reduce边平方边求和一气呵成floatsumthrust::transform_reduce(data.begin(),data.end(),[]__device__(floatx){returnx*x;},// 变换平方0.0f,// 初始值thrust::plusfloat()// 归约求和);// 1 次显存访问不需要临时数组执行过程读取 x → 算 x² → 累加到 sum ↑ ↑ ↑ 显存 寄存器 寄存器 每个数字的平方结果直接累加不存显存快速提速的三个法宝法宝 1合并连续的 transform只要看到连续的多个 transform就合并成一个// 看到这种连续操作transform(square);transform(sqrt);transform(add_one);// 立刻合并成transform([](floatx){returnsqrtf(x*x)1;});法宝 2用融合专用算子Thrust 提供了很多二合一算子分开做慢融合算子快功能transform reducetransform_reduce变换后求和/求最值transform inclusive_scantransform_inclusive_scan变换后前缀和// 求平方和的最优写法floatsum_sqthrust::transform_reduce(data.begin(),data.end(),[]__device__(floatx){returnx*x;},0.0f,thrust::plusfloat());法宝 3用惰性迭代器避免存储transform_iterator就像一个魔法管道数据流过时自动变换不占内存。// 创建一个平方管道autosquared_iterthrust::make_transform_iterator(data.begin(),[]__device__(floatx){returnx*x;});// 直接对平方后的数据求和无需临时数组floatsumthrust::reduce(squared_iter,squared_iterN);原理数据经过管道时才计算就像水龙头的水用的时候才流出来。什么时候融合能省多少计算规则假设有 N 个操作连续处理数据未融合的显存访问 N 读 N 写 2N 次 融合后的显存访问 1 读 1 写 2 次实际加速比连续操作数未融合访问融合后访问理论加速2 个4 次2 次2x3 个6 次2 次3x5 个10 次2 次5x操作越多融合收益越大什么时候不该融合融合虽好但有例外情况情况 1中间结果还要用// 平方结果要用好几次就别急着融合thrust::device_vectorfloatsquared(N);thrust::transform(data.begin(),data.end(),squared.begin(),square);// 平方结果被多次使用floatsumthrust::reduce(squared.begin(),squared.end());floatmax*thrust::max_element(squared.begin(),squared.end());// 这种情况保留中间结果反而更好比喻如果炒好的蛋要分给好几道菜用那就先炒好放着别每道菜都重新炒。情况 2融合太多导致寄存器不够// 融合几十个操作寄存器装不下thrust::transform(data.begin(),data.end(),result.begin(),[]__device__(floatx){floataop1(x);// 占用寄存器floatbop2(a);// 占用更多floatcop3(b);// ...// ... 30 个中间变量// 寄存器爆满反而变慢});比喻操作台就那么大同时处理太多食材反而手忙脚乱。一句话总结算子融合 让数据在寄存器里一条龙处理完减少往返显存的次数。核心记忆点显存慢寄存器快冰箱远操作台近中间结果留寄存器别存显存做菜别反复放冰箱连续操作合并成一个一气呵成实战口诀看到连续 transform→ 合并成一个 lambda需要变换求和→ 用transform_reduce想避免临时数组→ 用transform_iterator结语算子融合的本质其实非常简单能一次做完的事别分好几次做能放手边的东西别老往冰箱跑。理解了显存慢、寄存器快这个核心你就掌握了 GPU 性能优化最重要的思维方式。下次写 Thrust 代码时看到连续的操作记得问自己一句“这些操作能不能合并成一个”如果能恭喜你性能可能瞬间翻倍后记2026年8月15日于上海在claude opus 4.8辅助下完成。
返回列表