Бесплатная нормализация: как Meta прячет LayerNorm внутри тензорных ядер

Бесплатная нормализация: как Meta прячет LayerNorm внутри тензорных ядер

Нормализация — это налог, который платит каждая современная модель. LayerNorm и RMSNorm стоят в каждом блоке трансформера, стабилизируют обучение, ускоряют сходимость — и при этом пожирают до 20% тренировочного времени. Не потому, что они сложные, а потому что они бесполезны для тензорных ядер: это чистая работа с памятью, без единого умножения матриц.

Инженеры Meta из команды ADS Model Kernel Library опубликовали набор техник, которые позволяют «спрятать» нормализацию внутри GEMM и attention-ядер. Результат — до 90% скрытой задержки нормализации и до 35% ускорения всего attention-блока. Разбираемся, как это работает.

Почему нормализация — это проблема производительности

В архитектуре Kunlun — фундаментальной модели для рекламных систем Meta, на которой работает Generative Ads Model (GEM) — нормализация встроена практически в каждый ключевой компонент: Multi-Head Attention, Hierarchical Seed Pooling, PFFN на базе GDPA. В типичной LLM, где вычисления более сбалансированы, нормализация занимает около 10% общего времени. В рекламных моделях Meta, которые более memory-bound, доля вырастает до 20%. Это означает, что пятая часть тренировочного времени уходит на операции, которые не используют тензорные ядра вообще.

Корень проблемы — в различии стратегий тайлинга. Нормализация (LayerNorm, RMSNorm) по своей природе является операцией редукции: чтобы посчитать дисперсию или стандартное отклонение, нужно пройтись по всему внутреннему измерению. Каждый CTA (кооперативный потоковый массив, единица исполнения на GPU) загружает целые строки данных. А типичный GEMM тайлит вход по обоим измерениям — ни один тайл не охватывает всю строку целиком. Следовательно, GEMM не может «на лету» посчитать нормализацию: ему не хватает данных в рамках одного тайла.

Самое очевидное решение — растянуть тайл GEMM так, чтобы он охватил всю строку. Но тут возникают две проблемы. Первая: это отклонение от оптимальной стратегии тайлинга для чистого GEMM — страдает кэш, пайплайн, общая производительность матричного умножения. Вторая: на GPU Blackwell с 228 КБ разделяемой памяти, при bfloat16 и минимальных двух стадиях пайплайна, максимальный размер N ограничен 512. Для реальных моделей, где N может быть 4096, 8192 или больше, наивная физическая невозможна.

Lazy Pre-Norm: математический трюк, ломающий циклическую зависимость

Первая техника — Lazy Pre-Norm — решает проблему слияния pre-RMSNorm с последующим GEMM. Задача выглядит так: нужно вычислить C = rmsnorm(A) @ B, где rmsnorm(A) = A * rstd(A), а rstd — обратное стандартное отклонение строки.

Казалось бы, prologue fusion (слияние «до» GEMM) должен быть проще epilogue fusion — ведь каждый CTA в GEMM проходит через целые строки входного тензора A. Но есть циклическая зависимость: чтобы посчитать rstd, нужно просуммировать квадраты по всей строке, а для этого нужно пройти весь k-цикл. Но чтобы начать обрабатывать тайлы в k-цикле, уже нужен rstd! Классическая задача «курица и яйцо».

Разгадка приходит из простого математического наблюдения. В RMSNorm без elementwise affine (без обучаемых коэффициентов γ и β) нормализация — это просто построчное умножение на скаляр. А построчное умножение на скаляр коммутирует с матричным умножением:

(A * rstd) @ B = (A @ B) * rstd

Это значит, что умножение на rstd можно «лениво» отложить до конца k-цикла. Пока GEMM работает на тензорных ядрах, параллельно накапливается сумма квадратов элементов A. Когда k-цикл завершён, из суммы квадратов вычисляется rstd, и результат GEMM домножается на него. В результате нормализация становится обычным epilogue — дополнительной строкой кода после основного цикла, полностью скрытой за латентностью GEMM.

На практике Lazy Pre-Norm даёт 17–32% экономии задержки LayerNorm для маленьких форм (N = 64, 128). Но по мере роста N преимущество исчезает и превращается в регрессию: чем больше N, тем сильнее принудительный tile_n = N отклоняется от оптимального тайлинга GEMM.

У метода есть ограничения. Он не работает с elementwise affine (они умножают по столбцам, а не по строкам). Не работает с LayerNorm (там есть вычитание среднего, что нарушает свойство коммутативности). И обратное распространение получается нетривиальным — в прямом проходе нормализованный тензор нигде не материализуется, и его приходится реконструировать на лету.

Multi-CTA Norm Fusion: когда N слишком большой

Для больших N, где Lazy Pre-Norm не работает, инженеры Meta применили идею из Quack — нормализационных ядер на базе CTA-кластеров. Ключевая мысль: несколько CTA в одном кластере могут совместно обрабатывать одну строку данных, обмениваясь промежуточными результатами через распределяемую разделяемую память (DSMEM).

Любую нормализацию можно разложить на две части. Редукция — посчитать rstd (или mean + variance для LayerNorm) — требует обхода всей строки. Элементwise-применение — умножить каждый элемент на rstd — работает с отдельными элементами. Только редукция требует полного прохода по N, и именно её можно «разделить и завоевать» внутри CTA-кластера. Поскольку результат редукции мал (один скаляр на строку), накладные расходы на коммуникацию через DSMEM минимальны.

Этот multi-CTA алгоритм можно вставить прямо в epilogue GEMM. Каждый CTA в кластере вычисляет свой кусок строки результата GEMM, затем через DSMEM-коммуникацию совместно считает редукцию по всей строке, и каждый CTA применяет нормализацию к своему куску. Проблема «N слишком большой» решена — нагрузка распределяется между несколькими CTA.

Результат: до 90% скрытой задержки нормализации для типичных форматов рекламных моделей Meta. Причём это работает и для LayerNorm, и для RMSNorm, и с elementwise affine, и для произвольных N.

FlashNormAttention: мега-ядро для целого блока

Финальная и самая амбициозная техника — FlashNormAttention. Это слияние не просто нормализации с одним GEMM, а целого PFFN-блока архитектуры Kunlun: pre-LayerNorm, attention-ядро GDPA, residual connection, post-RMSNorm, второй residual — всё в одном мега-ядре.

Здесь сложность на порядок выше. Во-первых, GDPA — это multi-head attention, а нормализация применяется по всем головам одновременно. Значит, даже если один CTA видит полное измерение головы, для нормализации ему нужны данные от других голов. Это требует multi-CTA подхода, где CTA в кластере обрабатывают разные головы, но кооперируются для нормализации.

Во-вторых, memory-давление. Слияние такого объёма операций удваивает использование разделяемой памяти: нужно хранить промежуточные результаты (нормализованный Q, результат attention) для последующих residual-соединений. Инженеры применили три оптимизации: переиспользование буферов памяти для непересекающихся по времени данных, использование тензорной памяти (TMEM) и аккумулятора тензорных ядер для «бесплатного» сложения, и субтайлинг регистров — загрузка тензоров в регистры по чанкам, чтобы предотвратить spilling.

В-третьих, пайплайн. CUDA-ядерные операции нормализации блокируют тензорные ядра. Решение — тонкая warp-специализация: 8 варпов на основные вычисления (матричные умножения, активации), 4 варпа на RMSNorm, и отдельная пятая группа — на prologue-LayerNorm. Это позволяет перекрывать нормализацию предыдущей итерации с attention текущей.

Результат — до 35% ускорения всего attention-блока по сравнению с индуктор-компиляцией PyTorch. Бенчмарки проведены на NVIDIA B200 с 750 Вт power cap, в формате bfloat16, с типичными для рекламных систем Meta формами: batch 768, head_dim 128, K/V длиной 128, Q — разреженная последовательность с средней плотностью 0.5.

Какие DSL использовались

Вся работа выполнена на двух kernel DSL. TLX — набор расширений Triton с низкоуровневой поддержкой аппаратных особенностей GPU: управление warp-партициями, DSMEM-коммуникацией, TMA-операциями. Helion — высокоуровневый DSL с упором на скорость разработки, портируемость и исчерпывающий автотюнинг. Для нестандартных случаев вроде наивной fusion, где один из размеров тайла жёстко ограничен, Helion с его автотюнингом оказался незаменим.

Обе библиотеки открыты: код доступен в репозитории facebookresearch/ads_model_kernel_library на GitHub.

Что это значит для индустрии

Нормализация — это не «бесплатный» шаг предобработки, как её часто воспринимают. Это полноценный вычислительный этап, который в memory-bound моделях съедает пятую часть тренировочного времени. Работа Meta показывает, что «бесплатной» нормализацию можно сделать — но только если проектировать ядра с учётом конкретного hardware (CTA-кластеры Blackwell, DSMEM, TMEM) и конкретных паттернов доступа данных.

Для LLM-инфраструктуры это означает, что following-generation training frameworks должны будут учитывать fusion-оптимизации на уровне архитектуры ядра. Triton 3.x и Helion уже поддерживают нужные примитивы. Ожидается, что аналогичные техники появятся и у конкурентов — FlashAttention-3, cuDNN, vendor-specific kernels.

Особенно важен аспект warp-специализации в FlashNormAttention. Разделение варпов на специализированные группы — загрузку, матричные умножения, активации, epilogue и нормализацию — это паттерн, который хорошо масштабируется на другие fused-ядра. По сути, это ручная диспетчеризация ресурсов GPU: тензорные ядра работают с матрицами, CUDA-ядра обрабатывают нормализацию, и оба потока выполняются параллельно внутри одного CTA. Для исследователей это напоминание: бенчмарки «time-to-train» без учёта fusion-оптимизаций могут вводить в заблуждение. Двадцать процентов нормализации в Kunlun — это не приговор, а engineering challenge, который уже решён.

С точки зрения экосистемы, Meta публикует код в открытом доступе — это редкость для kernel-level оптимизаций production-моделей. Другие компании обычно держат такие наработки внутренними. Открытая реализация на Triton/Helion означает, что сообщество может адаптировать техники для своих архитектур — не только для рекламных моделей, но и для любых memory-bound трансформеров.

Часто задаваемые вопросы

Почему нельзя просто оптимизировать нормализацию отдельно?

Можно, и это делают — Quack показал эффективные standalone-ядра. Но даже идеально оптимизированная нормализация остаётся memory-bound операцией: она читает и пишет данные, не используя тензорные ядра. Единственный способ сделать её «бесплатной» — полностью перекрыть её латентность latency GEMM или attention, то есть выполнить те же вычисления, но внутри другого ядра, которое уже загружает те же данные.

Работает ли это для inference или только для training?

Техники применимы и для inference, но выгода может быть меньше. В inference модели часто уже квантизированы, batch-size меньше, а memory bandwidth используется эффективнее из-за меньших объёмов данных. Основные бенефициары — тренировочные кластеры с большими batch-size и длинными последовательностями, где нормализация становится узким местом.

Почему Lazy Pre-Norm не работает с elementwise affine?

Elementwise affine — это обучаемые параметры γ (scale) и β (shift), которые применяются по столбцам (по измерению embedding), а не по строкам. Математический трюк Lazy Pre-Norm основан на коммутативности построчного умножения с матричным умножением. Столбцовое умножение этим свойством не обладает — оно меняет результат в зависимости от того, до или после GEMM его применить.

Итог

Команда Meta показала три уровня слияния нормализации с вычислительными ядрами: наивную fusion для малых N (17–32%), Lazy Pre-Norm с математическим трюком коммутативности (до 32% для RMSNorm), и Multi-CTA Norm Fusion с CTA-кластерами и DSMEM (до 90% скрытой задержки). Финальный FlashNormAttention объединяет всё в одно мега-ядро, ускоряя целый attention-блок на 35%. Это не теоретическая работа — код открыт, бенчмарки сняты на production-железе B200, а техники интегрированы в крупнейшие рекламные модели Meta.

← Все записи