Учёные научили трансформеры отбирать контекст напрямую, а не через имитацию полного внимания

Учёные научили трансформеры отбирать контекст напрямую, а не через имитацию полного внимания

Разреживание внимания после предобучения, способ снизить квадратичную по длине последовательности стоимость механизма внимания в уже обученных трансформерах: вместо того чтобы сопоставлять каждый запрос со всем контекстом, для него выбирают лишь небольшой набор единиц контекста, токенов или их блоков. В существующих обучаемых методах для этого отбора обычно используют лёгкий модуль-селектор, который оценивает единицы контекста, а затем жёстко отбирает top-K лучших по этой оценке. Проблема в том, что такой жёсткий отбор, операция дискретная, и она блокирует градиент от целевой функции языкового моделирования, из-за чего селектор напрямую этой функцией не обучить. Поэтому такие методы обычно обучают селектор иначе, дистилляцией: его учат воспроизводить, как распределяется полное, «плотное» внимание исходной, неразреженной модели по слоям. Дистилляция действительно заставляет селектор ранжировать единицы контекста похоже на полное внимание, но это не то же самое, что ранжировать их по вкладу в итоговое предсказание при жёстко ограниченном бюджете внимания, то есть при ограниченном числе единиц контекста, которые модель вообще может учесть на один запрос. Из-за этого расхождения ограниченный бюджет рискует быть потрачен на единицы контекста, которые выглядят важными для полного внимания, но малополезны именно в условиях сокращённого бюджета.

Чтобы устранить это рассогласование, авторы предлагают SAS (Simple Attention Sparsification, «простое разреживание внимания»), механизм разреженного внимания с гейтами, который обучает ранжирование контекста сквозным образом, напрямую по целевой функции языкового моделирования, а не через дистилляцию. Ключевая идея, добавлять непрерывные, а не дискретные, оценки селектора прямо в логиты внимания уже во время обучения: тогда градиент от функции потерь доходит до селектора обычным обратным распространением ошибки, и жёсткий отбор больше не блокирует обучение.

Авторы называют три решения, которые оказались принципиальными для того, чтобы эта на вид простая схема заработала на практике. Во-первых, гейт нужно вносить внутрь softmax-функции внимания и именно в логарифмической форме, а не применять его отдельно после softmax. Во-вторых, гейты нужно нормализовать, чтобы сопоставлять оценку исторического, более раннего контекста с текущим блоком, который модель сохраняет всегда, вне зависимости от отбора. В-третьих, важно сохранять непрерывные оценки селектора, а не сводить их к жёсткому да/нет: тогда модель обучается ранжировать единицы контекста по относительной значимости, а не просто делить их на отобранные и неотобранные.

Для обучения на длинных последовательностях авторы реализовали экономичное по памяти вычислительное ядро на Triton, которое встраивает SAS в вычисления в стиле FlashAttention, алгоритма, который считает внимание по блокам, не материализуя целиком его квадратную матрицу в памяти.

В экспериментах на задачах рассуждения, понимания длинного контекста и в агентных сценариях SAS стабильно превосходит другие обучаемые методы разреженного внимания при разных бюджетах внимания, а при жёстко ограниченном бюджете разрыв особенно велик, по утверждению авторов, это показывает, что предложенное ранжирование контекста эффективнее отражает реальную полезность единиц контекста для итоговой задачи.

Ключевые факты

  • Разреживание внимания после предобучения снижает квадратичную стоимость механизма внимания трансформеров, выбирая для каждого запроса лишь небольшой набор единиц контекста (токенов или блоков), но обучаемые селекторы такого отбора обычно тренируют дистилляцией из полного внимания исходной модели, а не по реальной полезности единиц контекста при заданном бюджете внимания.
  • Авторы предлагают SAS, механизм разреженного внимания с гейтами, который добавляет непрерывные оценки селектора прямо в логиты внимания во время обучения: функция потерь языкового моделирования обучает селектор сквозным образом, обычным обратным распространением ошибки, минуя блокировку градиента при жёстком отборе top-K.
  • Три решения оказались принципиальными: гейт стоит внутри softmax-функции внимания в логарифмической форме; гейты нормализуются, чтобы сравнивать исторический контекст с всегда сохраняемым текущим блоком; оценки селектора остаются непрерывными, чтобы модель училась относительным приоритетам, а не жёсткому да/нет.
  • Для обучения на длинных последовательностях реализовано экономичное по памяти вычислительное ядро на Triton, встраивающее SAS в вычисления в стиле FlashAttention.
  • На задачах рассуждения, понимания длинного контекста и в агентных сценариях SAS стабильно превосходит другие обучаемые методы разреженного внимания при разных бюджетах внимания, особенно сильно при жёстко ограниченном.

Почему это важно

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

Кому это важно

Тем, кто обучает или дообучает большие языковые модели и заинтересован в снижении стоимости работы с длинным контекстом: предложенный механизм встраивается именно в обучение, а не только в инференс, и требует изменения того, как считаются логиты внимания. Инженерам инфраструктуры внимания и авторам библиотек для его эффективного вычисления, реализация SAS поверх вычислений в стиле FlashAttention с отдельным ядром на Triton даёт конкретный, воспроизводимый по описанию рецепт встраивания похожего механизма разреженного внимания с гейтами в существующий стек. Исследователям, которые уже строят или сравнивают обучаемые методы разреженного внимания, работа называет слабое место ранее распространённого подхода (обучение селектора дистилляцией из плотного внимания). Наконец, всем, кто использует модели на задачах рассуждения, работы с длинным контекстом или в агентных сценариях: именно на этих задачах, по данным авторов, SAS превосходит прежние подходы, особенно при жёстко ограниченном бюджете внимания.

Как это применить

Переносима сама схема: три названных авторами решения фактически образуют чек-лист для тех, кто реализует нечто подобное самостоятельно. Гейт нужно вносить внутрь softmax-функции внимания и именно в логарифмической форме, а не домножать на веса внимания отдельно после softmax. Гейты стоит нормализовать, чтобы историческая, более ранняя часть контекста была сопоставима по шкале с текущим блоком, который в схеме сохраняется всегда целиком. И главное, оценки селектора на протяжении обучения должны оставаться непрерывными: жёсткий отбор top-K применим на инференсе ради экономии, но именно непрерывность во время обучения даёт градиенту дойти до селектора и позволяет ему учиться относительным приоритетам единиц контекста, а не только бинарному отбору. Команде, которая работает с длинными последовательностями, для этого потребуется вычислительное ядро уровня FlashAttention, авторы закрыли эту часть отдельным экономичным по памяти ядром на Triton, что можно взять как ориентир при реализации похожей оптимизации в собственном стеке.

Можно ли доверять

Это страница на Hugging Face Papers, а не итоговая публикация в рецензируемом издании. Главное количественное заявление, что SAS «стабильно превосходит» другие обучаемые методы разреженного внимания и даёт «особенно большой выигрыш» при жёстком бюджете, исходит от самих авторов работы. При этом сама архитектурная идея описана конкретно и технически строго: указано точное место, где стоит гейт (внутри softmax, в логарифмической форме), названы все три архитектурных решения и то, зачем нужно каждое из них, и упомянута конкретная инженерная реализация, ядро на Triton под вычисления в стиле FlashAttention. Такая специфика на пустом месте обычно не пишется.

Риски и подводные камни

Главный риск для читателя, принять «стабильное превосходство» SAS и «особенно большой выигрыш» при жёстком бюджете за проверенный количественный факт, хотя вся оценка результатов, исключительно авторская. Есть и техническая цена решения: сам факт, что авторам потребовалось отдельное экономичное по памяти ядро на Triton, чтобы обучать SAS на длинных последовательностях, говорит о том, что затраты памяти при обучении, не тривиальная деталь, а то, что потребовало отдельной инженерной работы.

«Чтобы устранить это рассогласование, мы предлагаем SAS, механизм разреженного внимания с гейтами, который сквозным образом, вместе с функцией потерь языкового моделирования, оптимизирует ранжирование контекста.»

— авторы работы