FP8-обучение на AMD GPU: PyTorch ускорился на 13% из коробки
FP8-обучение обещает почти двукратное ускорение тренировки больших моделей, но на AMD Instinct оно долго давало не ускорение, а тихо испорченные градиенты. Причина банальна и неприятна: библиотеки PyTorch считали скейлы под формат NVIDIA, а железо AMD использует другую разновидность FP8. Ошибка не выбрасывала исключений и не падала с NaN. Модель просто училась хуже, и никто не понимал почему. Теперь это исправлено: AMD и Meta влили FP8-оптимизации напрямую в TorchAO и TorchTitan, и тренировка на MI300X и MI325X работает из коробки, без единого AMD-специфичного пакета.
Что такое FP8-обучение
FP8-обучение: это тренировка нейросети, где матричные умножения выполняются в 8-битном формате с плавающей точкой вместо привычного BF16. Каждый линейный слой делает три GEMM-операции: прямой проход, градиент по входу и градиент по весам. Квантизация этих операций с 16 до 8 бит ускоряет их за счёт более быстрых матричных ядер в кремнии. Это не про экономию памяти: в бенчмарке на Llama 3 8B пиковое потребление осталось почти тем же, около 39 ГБ. Выигрыш даёт именно скорость вычислений.
Результаты, которые команды AMD и Meta показали после апстрима, выглядят так. Плотная Llama 3 8B на восьми MI300X с rowwise FP8 дала +13,4% к throughput против BF16. Для MoE-архитектур картина сложнее: на DeepSeek-V3 671B квантизация сначала добавляла заметный оверхед, но после серии оптимизаций Triton-ядер удалось вернуть 89% этого разрыва. Отдельные ядра ускорились до 6,2 раза.
Тихая порча градиентов: проблема формата FNUZ
Самая интересная часть истории: не скорость, а корректность. AMD Instinct реализует вариант FP8 под названием FNUZ (Finite, No NaN, Unsigned Zero). Формат e4m3fnuz имеет максимальное значение 240 и, что важнее, не имеет кодировок для NaN и Inf. TorchAO изначально считал скейлы под NVIDIA-формат e4m3fn с другим максимумом. На железе AMD это означало, что тензоры масштабировались в диапазон, который железо физически не может представить.
Активации клиппились, градиенты портились, но поскольку в FNUZ нет NaN, переполнение нигде не всплывало как ошибка. Тренировка продолжалась, метрики деградировали, а причина оставалась невидимой. Выбор правильного формата здесь: не опция тюнинга, а требование корректности. Решением стало автоопределение платформы: TorchAO теперь сам выбирает правильный FP8 dtype и максимальное значение вместо хардкода под NVIDIA (PR #1142, #1150, #2225). Параллельно в TorchTitan исправили отчёт о пиковых FLOPS для MI300X, чтобы цифры MFU были честными (#920), и добавили платформенные loss-бейзлайны под арифметику FNUZ (#2156).
Скейлы: от тензора до микроблока
FP8-квантизация применяется с разной гранулярностью, и от этого зависит баланс скорости и точности. Самый быстрый и грубый вариант: один скейл на весь тензор (tensorwise). Точнее: скейл на строку (rowwise), ещё точнее: на тайл фиксированного размера (blockwise), а самый тонкий: MXFP8, где скейлы на группу упакованы прямо рядом с данными. TorchAO и TorchTitan поддерживают все четыре стратегии, и теперь каждая корректно работает с AMD-арифметикой. Отдельно влили blockwise-ядра для MI300 и MI350 (#3996).
MoE: почему DeepSeek-V3 ломал пайплайн
С плотными моделями всё относительно просто: у каждого линейного слоя одинаковая форма, скейлы равномерные. Mixture-of-Experts устроен иначе. Модели вроде DeepSeek-V3 и Llama 4 роутят каждый токен к подмножеству экспертов, и получаются батчи переменного размера, которые нужно прогонять через grouped GEMM. Здесь требуются скейлы на строку для активаций, скейлы на колонку для каждого эксперта в весах и тензор офсетов, направляющий строки к нужному эксперту.
FP8 grouped GEMM на ROCm включили через бэкенд Composable Kernel с правильным dtype и диспатчем под AMD (#3955). Токены роутятся по офсетам, квантизуются фьюзнутыми Triton-ядрами и уходят в GEMM одним запуском.
Три уровня оптимизации Triton-ядер
Когда корректность установлена, начинается скорость. Пайплайн квантизации в TorchAO превращает тензор в FP8 цепочкой шагов: посчитать per-row или per-column absmax, вывести скейл, применить его, отклиппить и привести к FP8. В наивной реализации каждый шаг: отдельный запуск ядра, и между шагами промежуточные тензоры материализуются в HBM. Для MoE-моделей с десятками тензоров весов экспертов на слой эти лишние ходки в память доминируют над всем оверхедом FP8. Восьмибитная математика дешёвая, а вот запуски ядер и транзиты в HBM вокруг неё: нет. Квантизация здесь memory-bound, поэтому оптимизации били по движению данных, а не по арифметике, на трёх уровнях.
Первый уровень: запускать меньше ядер. В backward-проходе паттерн .t().contiguous().t() форсировал полное копирование тензора через HBM только для смены layout под GEMM. Копии убрали (#3972), а многошаговую цепочку scale-and-cast сфьюзили в одиночные Triton-ядра (#4069). На DeepSeek-MoE-16B и восьми MI300X это дало 4,2-кратное ускорение backward-прохода.
В forward-проходе та же история: квантизация весов экспертов запускала пять generic-ядер на вызов, а при 24 вызовах на шаг это добавляло около 90 мс оверхеда на шаг. Всю цепочку заменили одним фьюзнутым ядром (#4311), которое параллелится и по экспертам, и по блокам выходной размерности. На восьми MI325X с DeepSeek-V3 671B это дало +17% к end-to-end throughput: с 5996 до 7027 токенов в секунду. Для сравнения, бейзлайн BF16 на том же железе: 7156 tok/s. То есть фьюжн вернул почти весь разрыв.
Второй уровень: заставить каждое оставшееся ядро эффективно двигать память. Ядро colwise-скейлов в backward делало некоалесцированные записи: соседние SIMD-лейны писали по адресам в K байт друг от друга, и каждая запись порождала отдельную транзакцию в память. Фикс: транспонировать выходной тайл через LDS (Local Data Share) перед сохранением, плюс фьюзнутый однопроходный вариант, убирающий лишнее чтение из HBM. Результат на MI300X для шейпов DeepSeek-V3: с 7290 до 1170 микросекунд на MoE-слой, то есть 6,2-кратное ускорение (#4113).
Третий уровень: убрать синхронизацию, которая железу не нужна. Атомарные операции Triton (atomic_add, atomic_max, atomic_min) по умолчанию используют acquire-release порядок памяти. На AMD GPU это вставляет дорогие memory fence до и после каждой атомики, хотя для коммутативных редукций они избыточны. Порядок переключили на relaxed на AMD (#3945) с проверкой torch.version.hip, чтобы поведение на NVIDIA не изменилось.
Что не сработало
Честная деталь, которую редко пишут в релизных заметках. Команда расширила пространство автотюна Triton для MoE FP8-ядер с одной до 8-16 конфигураций (#3952), ожидая, что более широкий поиск найдёт более быстрые размеры тайлов под wavefront-архитектуру AMD. Бенчмарки на шейпах Llama 4 и MI300X показали нулевое улучшение, а лишние конфиги раздули время компиляции первой итерации. Изменение откатили (#4024). Вывод, который стоит унести с собой: пространство автотюна должно формироваться ограничениями железа (размер wavefront, ёмкость LDS, регистровое давление), а не расширяться по умолчанию.
Почему апстрим, а не отдельная библиотека
Чтобы понять ценность произошедшего, стоит посмотреть, как жила FP8-оптимизация на AMD раньше. На PyTorch Conference 2025 команда AMD демонстрировала линейное масштабирование за пределы тысячи GPU на Instinct-кластерах, но весь FP8-стек жил в Primus-Turbo, собственной оптимизационной библиотеке AMD поверх тренировочных фреймворков. Такой подход типичен для вендоров железа: своя библиотека, свои ядра, свои рецепты. Проблема в том, что фрагментированный стек почти никто не использует. Команды берут стандартный TorchTitan, потому что он документирован, поддерживается сообществом и не требует изучения чужой экосистемы. Вендорская библиотека остаётся демкой для конференций.
AMD выбрала другой путь: вместо развития Primus-Turbo как отдельного продукта её инженеры вместе с командой Meta перенесли оптимизации в мейнлайн. Вклад разошёлся по двум репозиториям: в TorchAO ушли поддержка FP8 dtype под AMD, автоопределение формата и Triton-оптимизации ядер, в TorchTitan: исправления MFU, loss-бейзлайны и рецепты скейлинга. Это стратегически важный сигнал для всей экосистемы ROCm: AMD наконец играет вдолгую, инвестируя в стандартный стек, а не в параллельную вселенную.
Контекст здесь: конкуренция с NVIDIA за тренировочные кластеры. Аргумент против AMD традиционно звучал как «железо неплохое, но софт сырой». Каждый такой апстрим подрывает этот аргумент: FP8, grouped GEMM для MoE и корректные скейлы: это как раз те места, где «сырость софта» ощущалась физически, в виде испорченных градиентов и двузначных потерь throughput.
Что это значит на практике
Главный итог не в процентах, а в модели поставки. Раньше FP8 на AMD жил в Primus-Turbo, оптимизационной библиотеке AMD, которую надо было отдельно ставить и интегрировать. Теперь всё влито в мейнлайн pytorch/ao и pytorch/torchtitan: команды с Instinct-картами получают ускорение простым апгрейдом версий, без единого AMD-специфичного пакета. Это снижает порог входа для FP8-тренировки и делает AMD-кластеры заметно конкурентнее в экономике обучения: при сопоставимой цене железа +13-17% к throughput напрямую режут стоимость тренировочного прогона.
Дальше в работе MXFP8 grouped GEMM и квантизационные ядра для forward и backward на MI355X, результаты обещают в следующем посте. Пайплайн фьюжна (#3972, #4069, #4113, #4311) продолжит обрастать Triton-оптимизациями.
Часто задаваемые вопросы
FP8-обучение экономит память?
Нет, и это контринтуитивный момент. В бенчмарке Llama 3 8B пиковая память при rowwise FP8 осталась почти идентичной BF16, около 39 ГБ. Выигрыш в 13,4% throughput дают более быстрые FP8-матричные ядра, а не сокращение объёма данных.
Нужно ли ставить что-то специфичное для AMD?
Нет. Все описанные оптимизации слиты в мейнлайн TorchAO и TorchTitan. Достаточно обновить обе библиотеки: автоопределение платформы само выберет правильный FP8-формат, скейлы и диспатч ядер под ROCm.
Почему на MoE-моделях FP8 сначала замедлял тренировку?
MoE-модели требуют grouped GEMM с per-row и per-expert скейлами, а наивный пайплайн квантизации запускал пять ядер на вызов с промежуточными записями в HBM. При десятках экспертов на слой этот оверхед на движение данных съедал выигрыш от быстрой 8-битной арифметики.
Какая гранулярность скейлов лучше для начала?
Начните с rowwise: это рабочий баланс скорости и точности, проверенный на Llama 3 8B (+13,4% при честном весоградиентном рецепте, где GEMM обновления весов остаётся в BF16). Tensorwise быстрее, но грубее, blockwise и MXFP8 точнее, но требуют более новых ядер и подходят для чувствительных к качеству прогонов.
Итог
FP8-обучение на AMD Instinct перешло из категории «работает с патчами и молитвами» в категорию «обновите TorchAO и TorchTitan». Плотные модели получили +13,4% throughput, MoE-гиганты вроде DeepSeek-V3 вернули 89% квантизационного оверхеда, а история с форматом FNUZ: хорошее напоминание, что в низкоуровневой оптимизации корректность важнее скорости. Если у вас есть MI300X или MI325X и тренировка упирается в throughput: сейчас самое время попробовать rowwise FP8 на реальном ворклоаде.