Материал подготовлен автоматически по первоисточникам: ссылки на них — в конце статьи.
Точные вторые производные softmax attention научились считать на длинном контексте без хранения полной матрицы внимания. FlashBoB от команды George Mason University обработал контекст, на котором прежние точные реализации уже исчерпывали память, хотя препринт не рецензирован и все числа в нём получили сами авторы. Методы с обучением во время работы модели, градиентной памятью и оптимизацией второго порядка можно строить без замены точного внимания приближённым.
Зачем проходить назад через уже вычисленный градиент
Обычное обучение трансформера выполняет прямой проход и один обратный: сначала получает выход механизма внимания, затем считает градиенты для матриц запросов, ключей и значений. FlashAttention ускоряет оба этапа, потому что обрабатывает матрицу внимания блоками и не записывает её целиком в основную память GPU.
Этого недостаточно, когда градиент становится частью самого вычисления. Sophia-H оценивает кривизну функции потерь через произведения матрицы Гессе на вектор. Метаобучение дифференцирует внутренние шаги оптимизации, а обучение во время работы модели и градиентная память проводят обратный проход через обновление параметров.
В этих случаях нужен обратный проход по обратному проходу (backward-over-backward, BoB). Он принимает градиенты первого обратного прохода и вычисляет, как они меняются относительно исходных входов внимания и входящего градиента.
Стандартное автоматическое дифференцирование сохраняет для этого несколько промежуточных тензоров размером «длина последовательности на длину последовательности». Расход памяти растёт квадратично, поэтому длинный контекст быстро перестаёт помещаться на GPU. FlashAttention эту операцию не реализует: его схема памяти рассчитана только на прямой и первый обратный проходы.
Два скалярных состояния заменяют квадратичные тензоры
Второй обратный проход зависит одновременно от сумм по строкам и столбцам матрицы внимания. Если выполнять его напрямую, приходится либо несколько раз читать одни и те же данные, либо хранить полные промежуточные матрицы, либо разрешать множеству потоков одновременно обновлять общие результаты. Последний вариант использует FlashBack, но конкурирующие записи задерживают вычисление.
FlashBoB сводит глобальные зависимости внутри каждой строки к двум скалярным состояниям. Остальные величины зависят от них линейно: промежуточные суммы можно накопить заранее, а поправки применить после того, как строка обработана целиком. Благодаря этому естественная схема из нескольких обходов сокращается до двух.
В первом проходе алгоритм удерживает блок строк в быстрой памяти SRAM на кристалле и последовательно читает блоки ключей и значений. Вероятности внимания он заново вычисляет из сохранённых нормализаторов, накапливает строковые суммы, затем записывает два скалярных состояния и готовые градиенты для запросов и входящего градиента.
Во втором проходе ориентация меняется: в SRAM остаётся блок столбцов, а строки проходят через него потоком. Алгоритм повторно вычисляет нужные блоки вероятностей, читает сохранённые строковые состояния и завершает градиенты для ключей и значений. Полная матрица внимания ни на одном этапе не попадает в HBM, а каждый готовый блок результата записывается один раз.
Теоретическая оценка действует для принятой в FlashAttention модели, где оценки внимания повторно вычисляются, а SRAM достаточно велика для рабочих блоков. В этих условиях обмен с HBM достигает той же нижней границы, которая наследуется от точного прямого прохода внимания.
Изолированный тест использовал форму слоя GPT-2 Small, причинную маску, BF16 для входов и FP32 для суммирования на A100 с 80 ГБ памяти. Отдельный опыт проверял обучение GPT-2 Small на миллиарде токенов FineWeb с Sophia-H и восемью A6000. Реализация написана на Triton и охватывает стандартное многоголовое внимание.
Когда FlashBoB меняет инженерный план
В тесте отдельного слоя FlashBoB дошёл до последовательности в 262 тысячи токенов. Точные ядра на основе GradMem переставали помещаться уже к 16 тысячам, а относительно FlashBack новый алгоритм работал до 6,3 раза быстрее. Это не результат полного обучения на таком контексте: здесь измеряли только операцию второго обратного прохода через слой внимания.
В опыте с Sophia-H полный цикл обучения ускорился в 2,17 раза, а пиковое потребление памяти сократилось более чем вдвое. Кривые потерь двух вариантов Sophia-H практически совпали: новый путь изменил стоимость вычисления, но не траекторию оптимизации.
Для обычного обучения с первым обратным проходом и для инференса работа ничего не меняет: там уже достаточно FlashAttention. Она влияет на проекты, где точные вторые производные ранее заставляли сокращать контекст, переходить к приближённому вниманию или отказываться от внутреннего обучения модели.
FlashBoB заменяет только двойной обратный проход, поэтому обычные шаги могут сохранить оптимизированный путь внимания. Это важнее прямого сравнения с математической реализацией PyTorch: команда может добавить вычисления второго порядка, не замедляя все остальные шаги.
Текущую реализацию нельзя считать готовой заменой для любой архитектуры внимания. В работе стандартное многоголовое внимание проверено отдельно от полного обучения, а варианты с общими ключами и значениями, гибкими масками и страничным расположением кеша авторы относят к следующим расширениям. Для таких систем FlashBoB пока задаёт схему вычисления, но потребует отдельной интеграции и повторных замеров.
Источники
Похоже на вашу задачу?
Расскажите, что собираете. За полчаса разложим на этапы и назовём сроки — это бесплатно и ни к чему не обязывает.



