Математика глубокого обучения

Числа с плавающей точкой

fp32, fp16, bf16, fp8 — и почему у bf16 меньше точности, чем у fp16, но берут его

Шаг 78 из 117 · ~28 мин

Как устроено число

x=(1)s2ebias(1+f)x = (-1)^{\htmlData{k=sign}{s}} \cdot 2^{\htmlData{k=exp}{e - \text{bias}}} \cdot \big(1 + \htmlData{k=man}{f}\big)

Два бюджета, между которыми делятся биты:

  • задаёт диапазон — насколько большие и малые числа выражаются;
  • задаёт точность — сколько значащих цифр.

Ключевая величина — машинный эпсилон ε=2m\varepsilon = 2^{-m}, где mm — число битов мантиссы. Это расстояние от единицы до следующего представимого числа, и оно же — относительная погрешность любой операции.

Переключайте форматы. Обратите внимание, что все числа в таблице выведены из двух ширин, а не взяты из справочника.

раскладка битов

всего бит
32
эпсилон
1.19e-7
максимум
3.403e+38
мин. нормальное
1.18e-38
десятичных цифр
7.2
Стандарт по умолчанию: эпсилон 1.19e-7, около семи значащих десятичных цифр, диапазон до 3.4e38. Всё остальное в этом списке — попытки заплатить чем-то из этого за скорость и память.

Главный сюрприз: bf16 менее точен, чем fp16

Сравните две строки:

форматбиты мантиссыэпсилонмаксимум
fp1610109.81049.8 \cdot 10^{-4}6550465\,504
bf16777.81037.8 \cdot 10^{-3}3.410383.4 \cdot 10^{38}

У bf16 меньше битов мантиссы и в восемь раз больший эпсилон. По любой мере точности он хуже. И именно его используют для обучения.

Причина в том, что два вида поломок несимметричны:

  • потеря точности ухудшает шаг обучения на доли процента. SGD и без того работает с шумной оценкой градиента (блок 5, урок 060), и лишний шум в третьем знаке теряется в этом шуме;
  • переполнение или обнуление ломает обучение целиком. inf в градиенте портит все веса за один шаг; ноль означает, что слой не учится вовсе.

Диапазон fp16 (61056 \cdot 10^{-5} до 6550465\,504) для градиентов тесен, и с ним приходится делать loss scaling: умножать лосс на большую константу, чтобы поднять градиенты в представимую область, а потом делить обратно. Это работает, но требует динамического подбора масштаба и обработки переполнений. bf16 снимает проблему целиком — его диапазон совпадает с fp32, — и потому вытеснил fp16 в обучении, оставив ему инференс.

Эпсилон — это относительная величина

Частая ошибка — считать эпсилон абсолютной погрешностью. Он относительный: расстояние между соседними представимыми числами пропорционально самому числу.

околошаг в fp32
111.21071.2 \cdot 10^{-7}
10610^60.06\approx 0.06
101010^{10}103\approx 10^3

Отсюда практическое следствие: прибавление малого к большому теряется. В fp32 108+110^8 + 1 равно ровно 10810^8, потому что единица меньше шага сетки в этой области.

Это и есть механизм, из которого растут все проблемы следующего урока: накопление суммы из миллиона слагаемых, усреднение по большому батчу, вычитание близких величин. Ничего из этого не «неточно» — всё это точно настолько, насколько позволяет сетка, и именно поэтому порядок операций начинает иметь значение.

Что стоит помнить

форматгде применяют
fp32мастер-копия весов, накопители, чувствительные операции
bf16обучение: прямой и обратный проход
fp16инференс; обучение только с loss scaling
fp8инференс и, всё чаще, часть операций обучения

Заметьте строку про мастер-копию. В смешанной точности веса хранят в fp32, а вычисляют в bf16, потому что шаг оптимизатора часто меньше эпсилона: прибавление 10410^{-4} к весу порядка единицы в bf16 не изменит ничего. То есть fp32-копия нужна не для точности вычислений, а для того, чтобы обновления вообще накапливались.

Источники

Проверки

0 из 2
  1. Форматы и их компромиссы

    Отметьте все верные утверждения о числах с плавающей точкой.

  2. Вывести характеристики формата

    Реализуйте format_facts(exponent_bits, mantissa_bits) — по двум ширинам полей выведите всё остальное. Верните [total_bits, epsilon, max_value, min_normal, decimal_digits]:

    • total_bits = 1+e+m1 + e + m (знак, экспонента, мантисса);
    • смещение экспоненты bias=2e11\text{bias} = 2^{e-1} - 1;
    • epsilon = 2m2^{-m} — расстояние от единицы до следующего представимого числа;
    • max_value = 2emax(22m)2^{e_{\max}} \cdot (2 - 2^{-m}), где emax=(2e2)biase_{\max} = (2^{e} - 2) - \text{bias} (верхний код экспоненты зарезервирован под бесконечность и NaN);
    • min_normal = 21bias2^{1 - \text{bias}} — наименьшее нормальное число;
    • decimal_digits = log102m+1\log_{10} 2^{m+1} — значащих десятичных цифр с учётом неявного старшего бита.

    Проверить формулы можно на fp64 (e=11e = 11, m=52m = 52): они обязаны дать в точности sys.float_info.epsilon, .max и .min вашего языка. Если сошлось там, сойдётся и для остальных форматов.

    функция format_facts

    Загрузка редактора…

    Ctrl/⌘ + Enter