Иногда можно встретиться с проблемой при суммировании большого количества чисел с плавающей точкой. Особенно если числа разных порядков, эта ошибка может быть значительной для нашей дальнейшей логики.
Корень проблемы лежит в представлении чисел в памяти. По стандарту IEEE 754 число с плавающей точкой представляется знаком, экспонентой и мантиссой. Со знаком всё понятно (отрицательное или положительное число), экспонента, грубо говоря, показывает, где ставить точку в числе, мантисса представляет собой значащие цифры числа.
При обычном суммировании может не хватить точности мантиссы (при fp32 она занимает 24 бита).
Простой пример:
a = np.array([1e8, 1.0, 1.0, -1e8], dtype=np.float32)
print(np.sum(a))
Выведет:
0.0на единички не хватило 24 бит, потому что
log_2(10**8) — это примерно 26.5Тут на помощь приходит Kahan Algorithm:
def kahan_sum(arr):
total = np.float32(0.0)
compensation = np.float32(0.0)
for val in arr:
y = val - compensation
new_sum = total + y
compensation = (new_sum - total) - y
total = new_sum
return total
После этого сумма будет правильной:
2.0Почему это работает?
При обычном сложении ошибка округления теряется безвозвратно.
Алгоритм Кахана же сохраняет эту ошибку в переменной
compensation и вычитает её на следующем шаге.
