如何在保证精度的前提下加速Neumaier求和算法?
Neumaier求和的加速优化:精度与速度的平衡
Neumaier求和是Kahan求和的改进算法,用于精确对浮点数数组求和。
import numba as nb @nb.njit def neumaier_sum(arr): s = arr[0] c = 0.0 for i in range(1, len(arr)): t = s + arr[i] if abs(s) >= abs(arr[i]): c += (s - t) + arr[i] else: c += (arr[i] - t) + s s = t return s + c
该实现精度表现良好,但如果添加fastmath=True参数,速度至少能提升四倍。遗憾的是,fastmath允许对求和进行重括号优化(即结合律优化),这会破坏计算精度,因此无法直接使用该参数。
以下是不同求和方式的测试结果:
首先创建一个长度为1001的测试数组:
import numpy as np n = 10 ** 3 + 1 a = np.full(n, 0.01, dtype=np.float64) a[0] = 10**10 a[-1] = -10**10
我们引入math模块,用fsum获取理论正确结果,同时定义一个开启fastmath的Neumaier求和版本做对比:
# 使用fastmath的精度劣化版本 @nb.njit(fastmath=True) def neumaier_sum_fm(arr): s = arr[0] c = 0.0 for i in range(1, len(arr)): t = s + arr[i] if abs(s) >= abs(arr[i]): c += (s - t) + arr[i] else: c += (arr[i] - t) + s s = t return s + c
测试结果如下:
math.fsum : 9.99 nb_neumaier_sum : 9.99 nb_neumaier_sum_fm: 9.99001693725586
计时结果:
%timeit nb_neumaier_sum_fm(a) 350 ns ± 0.983 ns per loop (mean ± std. dev. of 7 runs, 1,000,000 loops each) %timeit nb_neumaier_sum(a) 1.5 µs ± 18.8 ns per loop (mean ± std. dev. of 7 runs, 1,000,000 loops each)
现提出技术问题:是否存在方法能够加速上述Neumaier求和代码,同时保持其当前的输出精度?
对应的汇编代码如下(注释由ChatGPT添加,若存在错误请指出):
movq 8(%rsp), %rcx # 将数组地址加载到rcx寄存器 movq 16(%rsp), %rax # 将数组长度加载到rax寄存器 vmovsd (%rcx), %xmm1 # 将数组第一个元素加载到xmm1(对应变量s = arr[0]) leaq -1(%rax), %rdx # 计算数组长度减1,存入rdx testq %rdx, %rdx # 测试rdx是否为0 jle .LBB0_1 # 如果rdx <= 0,跳转到.LBB0_1 movabsq $.LCPI0_0, %rsi # 加载绝对值掩码的常量地址到rsi movl $1, %edx # 初始化循环计数器i为1 vxorpd %xmm0, %xmm0, %xmm0 # 清空xmm0寄存器(对应变量c = 0.0) vmovapd %xmm1, %xmm3 # 将xmm1的值(s)复制到xmm3 vmovapd (%rsi), %xmm2 # 将绝对值掩码加载到xmm2 .p2align 4, 0x90 .LBB0_3: vmovsd (%rcx,%rdx,8), %xmm4 # 加载arr[i]到xmm4 vandpd %xmm2, %xmm3, %xmm5 # 计算abs(s),结果存入xmm5 incq %rdx # 循环计数器i自增1 vaddsd %xmm4, %xmm3, %xmm1 # 计算t = s + arr[i],结果存入xmm1 vandpd %xmm2, %xmm4, %xmm6 # 计算abs(arr[i]),结果存入xmm6 vcmpnlesd %xmm5, %xmm6, %xmm5 # 比较abs(s)和abs(arr[i]),结果存入xmm5 vsubsd %xmm1, %xmm3, %xmm6 # 计算(s - t),结果存入xmm6 vaddsd %xmm6, %xmm4, %xmm6 # 计算(s - t) + arr[i],结果存入xmm6 vsubsd %xmm1, %xmm4, %xmm4 # 计算(arr[i] - t),结果存入xmm4 vaddsd %xmm4, %xmm3, %xmm3 # 计算(arr[i] - t) + s,结果存入xmm3 vblendvpd %xmm5, %xmm3, %xmm6, %xmm3 # 根据比较结果选择对应值更新xmm3 vaddsd %xmm3, %xmm0, %xmm0 # 更新变量c vmovapd %xmm1, %xmm3 # 更新变量s(s = t) cmpq %rdx, %rax # 比较循环计数器和数组长度 jne .LBB0_3 # 如果不相等,继续循环 vaddsd %xmm1, %xmm0, %xmm0 # 最终计算s + c xorl %eax, %eax # 清空eax寄存器 vmovsd %xmm0, (%rdi) # 将结果存入目标地址 retq # 函数返回
内容的提问来源于stack exchange,提问作者Simd
相关产品推荐
相关产品推荐

