ARTICLE DETAIL

资讯详情

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

浮点运算工程实践:可复现性、误差控制与混合精度优化

浮点运算工程实践:可复现性、误差控制与混合精度优化 如果前面的七篇你都跟下来了我相信你对float和double的底细已经比大多数同事清楚符号位、指数、尾数、舍入模式、特殊值、ulp这些概念现在应该能脱口而出。但真到工程里还是会碰到很多“纸面上讲不通”的问题为什么同一套代码换个编译器结果就跑偏了为什么加了几个线程结果就开始跳为什么代码看起来没变升级硬件后二进制就不一样了这些问题的根源往往不是数学公式错了而是浮点运算的环境、舍入顺序和中间精度在暗地里发生变化。这一篇继续聊浮点运算但聊的是更工程向的部分可复现性、误差可控性和性能优化的取舍。适合正在调试数值代码、做并行计算或者打算用低精度紧缩成本的工程师参考。我分享的都是一些可以立刻用上的检查方法和避坑套路有的甚至算不上优雅但管用。1. 可复现性同一份代码凭什么两次运行结果不一样1.1 浮点环境不是一个常量先说一个很容易被忽略的事实IEEE 754 只约束了每一次运算的数学结果但没有约束你的源代码最终被翻译成哪几条 CPU 指令。于是同样一行c a * b d;在不同编译选项下可能有两种行为。如果编译器把目标平台支持的 FMA 指令用上了a*b的 53 位完整乘积会先保留下来再加上d最后只做一次舍入。如果不支持 FMA则会先算a*b舍入成一次 double再与 d 相加再做第二次舍入。两种路径都可能满足 IEEE 754 的“每步正确舍入”但最终结果可以在最后 1 个 ulp 上不一样。注意这里谁都不算错只是舍入发生的位置不同。再加上 x86 平台的历史包袱早期 x87 指令使用 80 位扩展精度寄存器做中间计算。如果你的程序把中间结果留在寄存器里继续运算它的精度其实比 double 高一旦被编译器溢出到内存又会被截断回 64 位。同一个表达式寄存器分配不同结果就不同。现代 64 位编译器基本都用 SSE 做标量浮点但老代码或者一些嵌入式交叉编译器仍然可能踩这个坑。所以想要可复现第一步不是找 bug而是把“浮点环境”固定下来统一编译器版本、统一优化档、统一 math 库、统一 CPU 指令集目标。如果一个项目允许用户调整优化选项请把浮点结果的可复现性作为调优时的一个显式测试项。1.2 并行归约里的顺序陷阱比编译环境更隐蔽的是并行归约。假设你要算 4 个数a, b, c, d的和串行顺序自然是((ab)c)d。一旦拆到多个线程线程 1 算ab线程 2 算cd最后合并出(ab)(cd)。在浮点里这两个排序不保证相等。道理很简单浮点加法不满足结合律。从数学上123怎么排都一样但从舍入角度看(0.10.2)0.3和0.1(0.20.3)就会相差约 2^-54 量级这不是 bug而是舍入顺序的必然产物。我踩过的坑出现在 OpenMP 归约上。代码里写#pragma omp parallel for reduction(:sum)机器核数一多每次跑出来的sum在小数点后十六位开始飘。数据本身只有几十万条绝对误差看起来很小但下游模块只要做一个减法误差就可能被放大。后来我把归约改成固定分组每个线程负责 1024 个元素部分和被存进固定长度数组再由主线程按固定顺序累加结果就完全可复现了。GPU 上更严重。atomicAdd的合并顺序由硬件调度决定同一个 kernel 跑两次都可能不同。如果你在 GPU 上做科学计算并且需要可复现要么避免 atomic要么把结果先做到每个 block 的私有数组最后用固定顺序二次归约。还有一个常规技巧是提升累加精度部分和用误差补偿下一节会展开但前提依然是顺序固定。顺序不固定任何补偿都很难救。2. 舍入误差从“知道会错”到“知道它怎么错”2.1 先看问题的条件数别急着改代码面对一个数值结果不对的 bug很多人的第一反应是换更高精度double 不行换 long doublelong double 不行换任意精度。这个方向不一定错但更高效的做法是先判断问题本身是不是“病态”的。数学上如果输入变化一点点输出就变化很大那么不管什么精度的浮点只要输入有舍入误差结果都不可靠。有一个简单工具叫条件数。对函数 f(x) 做近似条件数可以写成 |x f(x) / f(x)|。当条件数接近 1 的时候输入舍入误差不会放大太多当条件数远大于 1 时结果对误差极其敏感。最经典的例子就是减法抵消计算sqrt(x1) - sqrt(x)当 x 很大时两个平方根几乎相等相减后所有高位抵消只留下一点点低位噪声。此时条件数差不多等于 sqrt(x)/2属于病态表达式。但“病态”不等于“没法算”只要改写表达式就可以把sqrt(x1)-sqrt(x)改写成1/(sqrt(x1)sqrt(x))两边分母都是正数相加不再有灾难性抵消数值稳定性一下子就好了。这种代数重排不改变数学结果却彻底改变舍入误差的传播。所以定位浮点误差的第一步是问你算的是稳定表达式吗如果问题本身就病态再高的精度也只是把灾难推迟几级。2.2 用ULP当误差标尺而不是绝对误差第二个常见误区是拿绝对误差评价结果。0.1和0.100000000000000006差 6e-18看起来极小但如果你在计算 1e30 量级的数它的 ulp 可能已经是 1e14 了误差再小也救不会高级比特。与其说“误差小于 1e-15”不如说“误差小于几个 ulp”。C99 里给出nextafter可以方便地量出两个相邻可表示浮点数之间的距离#include math.h #include stdio.h double ulp_size(double x) { if (isnan(x)) return 0.0; return fabs(nextafter(x, HUGE_VAL) - x); }实际工程中不一定要逐个数 ulp可以用nextafter(x, HUGE_VAL) - x直接得到 x 这一点上两个相邻 double 的距离如果要比较两个计算结果直接把它们用memcpy转成uint64_t相减也能得到相差几个二进制最小步长的粗略度量。误差审查报告里写“误差 2 ulp”比写“误差 3e-16”有价值得多因为后者无法让你判断这个误差在大数和小数场景下的影响。遇到跨平台、跨精度的问题把所有数值结果打印成%a十六进制浮点用 ulp 或 bit pattern 差异而不是十进制小数来做回归判断能少掉一半的发际线。2.3 Kahan求和能救场但别指望它逆天然后是求和问题。最朴素的累加sum x[i]在数组长度很长、元素量级差别很大时误差会随 n 增长。经典的改进是 Kahan 补偿求和它在维护sum的同时记下每一步舍入丢掉的小尾巴cdouble kahan_sum(const double *x, size_t n) { double sum 0.0, c 0.0; for (size_t i 0; i n; i) { double y x[i] - c; double t sum y; c (t - sum) - y; sum t; } return sum; }为什么有效t sum y这一步的舍入损失会体现在(t - sum)与y的差里后一行正好把这个差提取出来下一轮再减回到y。实际测下来对 1e7 个随机 double 求和直接累加的相对误差能到 1e-13 量级Kahan 通常能压到 1e-15 附近。但注意两个前提仍然要求遍历顺序固定如果数组里有大量接近 1e308 的宏大量级Kahan 也会在加法本身溢出时失效。更彻底的办法是使用多精度累加比如 Shewchuk 的 expansion 求和或者用long double做中间累加器。long double在 x86 上 80 位精度能减缓但不保证可移植。在需要极高精度、又不想引入第三方库的时候我一般先用 Kahan 把数量级问题解决再加一个粗粒度的分段求和效果已经足够工程使用。3. 混合精度该省则省但边界要划清楚3.1 三种主流格式的脾气一张表看清楚性能优化的另一条路是降低精度。但低精度不是随便把double改成float更不是大家都在用的FP16就真的只有float一半精度。先看常见格式的量级表格式指数位尾数位含隐含位最小正规数最大有限数约十进制有效位double11532.2e-3081.8e30815~17float8241.2e-383.4e386~9FP165116.1e-5655042~3BF16881.2e-383.4e382~3注意两件事FP16 的最大值只有 65504在深度学习中一个大矩阵乘的中间值就可能超过它BF16 虽然尾数少但指数范围和 float 一样所以不会像 FP16 那样轻易上溢或下溢。因此从 double 降到 float你丢的是小数位从 float 降到 FP16你还可能丢范围。降精度之前先做一次范围审计你的数据会碰到1e30吗最小梯度会小于1e-20吗如果是FP16 可能直接变成 0 或 inf而不是“少几位精度”。这就是为什么工程上谈低精度总要先说范围再说有效位。3.2 低精度训练里的三项铁律目前混合精度训练在深度学习里已经是标配但很多人只把它当做一个开关并不懂背后的约束。我总结三条铁律缺一不可。第一主权重必须保持高精度。模型权重在每次更新时可能只变化1e-4量级这个增量远小于 FP16 在权重数值附近的 ulp。假如用 FP16 存权重连续几次梯度更新都会被舍入吃掉模型直接不收敛。所以框架里普遍维护一份 FP32 的主权重副本前向和反向计算时才把它 cast 到 FP16。第二Loss 要放大。FP16 的最小正规数约 6.1e-5梯度经过链式法则后会变得很小很容易下溢成 0。解决思路是在反向传播前把 loss 乘一个 scale常见是 1024 或 4096梯度也会被同等放大等更新主权重前再除回来。这个 scale 什么时候调大、什么时候调小就是第三点。第三动态损失缩放。如果某个 batch 算出来的梯度在放大后仍然出现 inf 或 NaN说明 scale 太大就把它缩小如果连续一段时间没有异常就试着调大。你不需要手工干预PyTorch 的GradScaler已经把这套逻辑封装好了。但理解原理之后遇到精度问题你就知道该去看梯度统计而不是盲目换优化器。3.3 更低精度当预处理器用高精度残差拉回混合精度不只存在于深度学习。数值线性代数里有个经典套路叫迭代精化先用低精度做一次快速的近似分解得到近似解x0再回到高精度计算残差r b - A x0然后解低精度修正方程A δ r最后更新x x0 δ。两步以后主解的精度基本恢复到高精度水平但大量重计算发生在低精度路径上速度接近低精度。这个套路能成功的原因很简单低精度主要承担大量低价值计算高精度残差负责捕捉被丢掉的部分。和 Kahan 的补偿思想本质上一样——把每一次舍入丢掉的“小尾巴”记录下来回到主线时再加回来。不过使用迭代精化要控制两个细节残差必须用高精度计算否则修正本身也带噪声低精度求解器必须是稳定算法不能病态到第一步就崩掉。如果你在写自己的混合精度模块可以先按这个思路试主流程低精度误差项高精度在接口层保留一个“是否允许低精度路径”的开关用一组固定测试向量做回归。一旦精度不合格还能退回纯高精度。4. 性能优化中的浮点原则别和编译器对着干4.1 编译器不是不改而是默认不敢改很多人一谈起浮点性能第一反应是“开-ffast-math就好了”。但在开之前先弄清楚编译器为什么默认不这么做。IEEE 754 允许 NaN、允许有符号零、允许正负无穷这些特殊值对优化形成了很多约束。比如x - x在数学上等于 0但当 x 是 NaN、正负无穷或±0时IEEE 语义下不能安全替换成 0x * 0也不能简化成 0因为 x 可能是 inf结果是 NaN或 x 是 NaN结果仍是 NaN。编译器如果自作聪明就会破坏程序对这些特殊值的处理方式。所以默认情况下编译器遵守“可观察浮点行为不变”的规则很多看似一步到位的优化都不合法。这时候你应该做的是用更细粒度的编译选项而不是直接上-ffast-math。GCC 有一些选项单独放宽某一条规则比如-fno-signed-zeros、-fno-trapping-math、-ffinite-math-only。如果你能保证自己的数据里不出现 NaN 和 inf-ffinite-math-only的收益通常比整个-ffast-math来得可控也更容易回归。我个人对-ffast-math的态度是小程序、无特殊值、有严格测试可以开大型数值库、涉及数学函数、需要跨平台复现尽量别开。开着它确实能快但查错的成本常常远超省下的几个机时。4.2 近似指令要用但不能无脑用硬件层提供了一些“反常规”的快速指令最有名的就是近似倒数平方根。游戏和图形社区经常讨论rsqrtss这类指令它能用几条指令给出接近1/sqrt(x)的估计初始相对误差大约在千分之一级别。如果你对精度要求不高直接用没问题如果单精度就够用可以再补一步牛顿迭代float quick_inverse_sqrt(float x) { float y hardware_rsqrt_approx(x); // 比如 _mm_rsqrt_ss return y * (1.5f - 0.5f * x * y * y); }一步迭代后误差能压到接近 IEEE float 的极限。注意两个前提x 是正数不是 0 也不是 inf/nanhardware_rsqrt_approx在不同架构的精度可能差一两倍仍要按最差情况做测试。另外1/sqrt(x)在遇到 x 下溢时硬件近似和公式恢复的路径也不同所以具体用之前要跑一遍全范围的边界测试。相似地exp、log也可以用多项式或查表近似但问题在于它们的动态范围很宽稍微偏掉一位指数结果就变形。我一般建议先用标准库函数用 profiler 确定热点再考虑近似近似前写下误差预算比如“最终结果相对误差小于 1e-4 就算达标”然后用随机输入做统计确保 99.9% 的样本在预算内。否则为了一个看起来更快的指令往往会把整个应用的正确性拖下水。4.3 性能优化前先量化优化后测“误差预算”性能优化里有个很反直觉的规律你越觉得一条指令慢越可能不是它的锅。我的习惯是先记录优化前的浮点结果作为基准再开 profiler 看热点找不到热点就不碰代码。真要优化把改动控制在最小范围。比如只把内层循环里double的sqrt换成带一次牛顿的近似版本同时保留外面所有高精度计算而不是把整个程序切到-ffast-math再花一周追查哪里出现了负零。每次改动后跑同一组随机回归用例记录每个中间变量的%a位模式和基准对比。下面这张表是我常用的选择框架供参考应用场景可接受误差推荐做法图形渲染、游戏物理千分之一到百分之一可用低精度/近似指令但做边界钳制深度学习训练/推理相对误差 1% 内FP16/BF16 混合精度主权重保持高精度科学计算/结构分析接近双精度或更严不开 fast-math保留严格 IEEE日志统计、监控告警不敏感可开快速近似但尽量先 profiler表格不是教条核心是“误差预算”你先定义多少误差可以接受再决定用什么手段。优化完成时如果误差在预算内就算数。不在就回退。5. 浮点调试三板斧异常、位模式和黄金参考5.1 打开异常铁门让程序在第一个错误处停下来调试浮点代码最忌讳的事是等程序算出 NaN 再回头找。默认状态下CPU 遇到非法操作、除零、上溢不会中断只是把结果写成一个特殊值继续往后跑。后面可能还要算很久才会发现“结果不对”但那时源头已经淹没了。把浮点异常打开程序会在第一处非法操作就崩掉配合调试器可以直接看调用栈。Linux 上可以用feenableexceptglibc 扩展#include fenv.h void enable_fp_exceptions(void) { feenableexcept(FE_INVALID | FE_DIVBYZERO | FE_OVERFLOW); }Windows 上则用_controlfp_s。这一类接口都不是 C 标准不同平台名字不一样但思路相同把异常掩码从默认的忽略改成抛错。打开异常以后程序可能会在0.0/0.0等原本合法的场景例如用户输入导致崩掉。所以测试阶段打开、生产环境恢复默认或者只在数值回归用例里打开。另外如果你开了-ffast-math它常常会顺手把浮点异常关了调试时记得关掉快速数学让异常机制正常工作。5.2 用十六进制位模式记录浮点现场遇到跨平台结果不一致很多人用printf(%.17f, x)对比这其实不够。十进制打印会经历一次“double 到十进制字符串”的转换不同标准库实现的舍入可能不同打印出来的最后几位不一定能反映真实的二进制差异。更可靠的方式是用%aprintf(x %a\n, x); // 输出形如 0x1.999999999999ap-4或者直接把 double 的 64 位模式打印成十六进制#include stdint.h #include string.h uint64_t double_bits(double x) { uint64_t bits; memcpy(bits, x, sizeof(bits)); return bits; }在日志里同时记下%a和 bit pattern两份日志一比对就能精确知道两个环境差了几个 ulp以及差别发生在哪一步。我自己做跨平台回归时会写一个小脚本把两组输出文件按行读出来用nextafter或 bit 差值统计每一行的误差而不是用 grep 找哪一个数字变了。这样效率高得多。5.3 把踩过的坑整理成一张速查表最后分享一张我自己整理的“浮点常见症状速查表”每次排查都从这张表开始症状常见原因排查/处理方向结果每次运行不一样并行归约顺序、硬件调度、乱序合并固定归约顺序避免 atomicAdd部分和高精度累加两个平台/编译结果差 1 ulpFMA contraction、数学库实现不同、x87 超精度统一编译器和 libm用%a/bit 对比必要时显式关闭 FMA小数值相加变成 0FP16/浮点下溢或累加顺序遇到大数使用更大精度的累加器loss scaling分块求和大数旁边加小数结果没变化当前量的 ulp 大于小数不要执着于浮点加法先用 Kahan/排序/分段开了-ffast-math后行为大变编译器放宽了 NaN、inf、符号零和结合性缩小优化范围改用细粒度选项重新跑边界用例算出的 NaN 出现得很晚默认异常屏蔽错误被传播到远处打开 FP 异常在源头捕获非法操作这张表是我常年迭代下来的版本不一定覆盖所有领域但最常见的几类都能对号入座。给自己项目做浮点回归时把遇到的新坑补充进去慢慢就会变成团队里最值钱的一份内部文档。从我个人的经验看浮点问题最麻烦的地方不是“怎么算”而是“我根本没有意识到它已经发生”。所以我现在的习惯是每个数值模块先写一组self_check禁掉快速数学打开异常再跑随机用例跨平台交付前固定编译器和 CPU 指令集用 bit pattern 做一次最终对比。这套流程不花多少时间但省掉的半夜排查时间实在太多。
返回列表