🔬 Метод
В основе эффективных техник (Training LLMs with MXFP4, Quartet, FP4 All the Way)
обучения неизбежно фигурирует стохастическое округление (SR) градиентов на обратном проходе. При стохастическом округлении тензор квантизуется не к ближайшему значению, а с вероятностью, зависящей от расстояния до ближайших точек на решетке вверх или вниз. Стохастическое округление обладает свойством несмещенности - матожидание сходится к истинному среднему. И это свойство критично для градиентов. Без него в низкой битности обучения нормально не сходятся.
Однако, оно менее точно, чем округление к ближайшему (RtN) и заметно повышает среднеквадратичную ошибку. Можно ли добиться среднеквадратичной ошибки почти как у RtN, при этом сохранив несмещенность?
И авторы статьи предложили способ этого добиться.
Несмещенности добиваются за счет двух приемов
1️⃣ Адаптируют метод EDEN из распределенной отпимизации, где применяется случайное Адамарово вращение и домножение на фактор, равный отношению скалярного произведения исходного вектора на скалярное произведение квантизованного повернутого и неквантизованного повернутого.
2️⃣ Скейлы квантизации стохастически квантизуются в FP8. Ошибка квантизации в FP8 куда меньше, чем у FP4.
Дополнительно на прямом проходе используют свежий трюк из статьи, где в зависимости от того, меньше ли ошибка квантизации на полной FP4 сетке (от -6, 6) или урезанной от -(4,4) выбирают первую или вторую.
MS-EDEN в итоге достигает вдвое меньшей ошибки по сравнению с SR.
Так же сравнивают блочную квантизацию 16x16 и линейные группы 1x16. Используют вторые, так как они дают меньшую ошибку.
🧪 Эксперименты
Метод валидируют сначала в постановке из Quartet I, обучая семейство моделей Llama-2 архитектуры. В качестве бейзлайнов берут рецепт от Nvidia , TetraJet-v2, и 4/6. Quartet II достигает наименьшего (до 20% по сравнению с лучшим бейзлайном) разрыва по лоссу в сравнении с bf16 обучением.
Затем обучают Nanochat с использованием всовременных наработок по обучению LLM - мюон оптимизатор, WSD расписание learning rate, ReLU2 активаций. Учат на FinewebEdu в режиме Шиншилла-оптимального скейлинга. Для наночата справедливы те же выводы о превосходстве Quartet II измеряемом в bits-per-byte.
И что приятно, все это приходит в связке с солидным ускорением. Матричные умножения ускоряются более чем в 4 раза в сравнении с bf16, а end-2-end throughput на обучении до 2.7x.
💡 Выводы
Данный результат дает дополнительную мотивацию к переходу на низкобитное обучение больших языковых моделей. Полагаю, что данный рецепт будет перенят Нвидией и иными крупными игроками.
Post #640
2.02K