TGViewer
КПД КПД @quant_prune_distill · 3.49K subscribers
Post #751 1.66K
🔬 Метод

Концептуально метод довольно несложный. Возникают две ветки вычислений:

🤙 Легковесная индексирующая ветка. В ней на каждую 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.
  • ❤ 7
More from @quant_prune_distill
  1. Oct 4, 2026"Горячие" эксперты
  2. Oct 4, 2026photo post
  3. Oct 1, 2026🛠 Метод Типичный scaling law имеет вид: L(N, D) = A N^α + B D^β + c 🔄 Скейлинг по рекурс…
  4. Oct 1, 2026Scaling Laws for Looped Mixture of Experts 📄 Статья Есть MoE, которые как-то скейлятся (п…
  5. Sep 29, 2026⚙️ Метод На префилле имеем дело с тяжёлыми матричными умножениями, поэтому для ускорения ц…
  6. Sep 29, 2026🧩 Disaggregated Quantization: Specializing LLM Prefill and Decode 📄 Статья Обычно для пр…
Threads Profile ViewerView any public Threads profile without an account.Open ThreadLook →Writing with AI? Make it sound human.Metric37 rewrites AI drafts so they read naturally. Free AI detector, 1,500 words free.Try Metric37 →