🔬 Метод
Концептуально метод довольно несложный. Возникают две ветки вычислений:
🤙 Легковесная индексирующая ветка. В ней на каждую query-группу в исходном внимании приходится одна query-голова и одна единственная key-голова на весь слой. Сначала считается pre-softmax скор, а затем отбирается максимальное значение в рамках некоего разбиения на блоки. Затем берутся topK блоков с наибольшим найденным значением максимума.
🤙 Главная ветка. Обычное блочно-разреженное внимание на блоках, которые вернула индексирующая ветка.
Процедура обучения устроена следующим образом:
👉Функция потерь — KL-дивергенция между выходами индексной ветки и усредненными выходами основной ветки. Хотим, чтобы выбранная sparsity в индексной ветке совпадала с главной. Через вероятности в главной ветке делается stopgrad, чтобы не коллапсировало.
👉 На входы индексной ветки тоже накладывается stopgrad, чтобы оптимизировались только проекции индексной ветки и градиент не тек ниже по модели.
👉 Разогрев индексной ветки. Сначала некоторое время гоняют полный attention в обеих ветках, а затем включают topK в индексной ветке.
👉 Последний блок (скользящее окно токенов) добавляют всегда.
С Minimax Sparse Attention (MSA) можно обучать как с самого начала, так и пострейнить.
Еще они наворотили topK-кернел, который быстрее торчового и быстрее TileLang-референса. И еще там персистентная сетка, учитывающая неоднородную и динамическую загрузку в процессе вычислений.
🧪 Эксперименты
Эксперименты проводят на 109B MoE с 6B активными параметрами (учат на 3T токенов суммарно). Есть три постановки:
🔅 Full Attention бейзлайн
🔅 MSA-PT (учат MSA с нуля)
🔅 MSA-CPT (2.6T токенов с Full Attention / 400B с MSA)
MSA учится довольно стабильно, лосс почти не отходит от Full Attention, без спайков и скачков.
По метрикам все 3 варианта довольно близки, без явного победителя, при этом в MSA внимание проводится всего на 2048 токенов (16 блоков по 128 токенов). Для работы на длинном контексте дополнительно дообучают модель, и там она на Helmet/RULER примерно на уровне MSA.
На контексте в 1М токенов предложенный подход дает 28х теоретического ускорения против полного внимания, но на практике выходит 14х на префилле и 7х на декоде из-за дополнительных оверхедов с запуском второй ветки вычислений, topK, индексаций.
💡 Выводы
Подход, безусловно, интересный, но не хватает сравнения с линейными вниманиями (в идеале гибридами полного и линейного внимания), а также другими бейзлайнами по Sparse Attention (с тем же Big Bird, DeepSpeed Sparse Attention, XAttention, DeepSeek-V4 разреженными вниманиями и многими другими). Также непонятно, как оно себя ведет на context-intensive задачах. А еще ускорения почему-то репортятся на H800, а кернелы выложены только под Blackwell.
Post #751
1.66K
- ❤ 7