Небольшой светящийся тайл движется по пустой матричной рамке, оставляя лишь компактные следы текущего состояния и выход со значением потерь.

Слитное ядро прямого прохода для линейной cross-entropy на GPU Apple

Denis Ineshin · 2026-07-19

Экономия памяти, пределы тайлинга и уроки несработавших моделей производительности

При обучении языковой модели ради одного лишь подсчёта cross-entropy нередко выделяется огромная матрица логитов. Для N токенов и словаря размера V её форма — (N, V). Но функции потерь из каждой строки нужны лишь log-sum-exp и один логит целевого токена. Мы написали слитное (fused) MLX/Metal-ядро прямого прохода, которое вычисляет эти величины по тайлам и ни разу не хранит матрицу целиком. При N=8192, V=151936 и скрытой размерности D=4096 слитный прямой проход занял 2.105 s против 1.286 s у материализованного MLX. Публичный прогон по памяти также зафиксировал заметно меньший пик для слитного пути. Однако из-за неодинаковых базовых уровней после прогрева корректно посчитать отношение памяти на один вызов нельзя; раздел 3 приводит полные измерения и объясняет это ограничение.

Оптимизация принесла два менее очевидных результата. Во-первых, принудительный вызов mx.eval после каждого чанка на чистом MLX увеличивал расход памяти: MLX удерживал трассируемые промежуточные значения. Во-вторых, модель повторного использования загрузок верно предсказала ускорение для одного регистрового тайла и тут же ошиблась на следующем. Тайлы побольше работали медленнее. Последующие контрольные эксперименты отбросили несколько объяснений, но не смогли выделить единственную причину. Оставшиеся данные указывают на компромисс между повторным использованием загрузок на уровне исходного кода и расходом ресурсов в скомпилированном ядре — но ни то, ни другое мы не измеряли напрямую на уровне железа.

Эта статья описывает реализацию и эти измерения. Нового алгоритма cross-entropy она не предлагает. Линейная cross-entropy без материализации логитов и потоковые редукции softmax уже описаны в предыдущих работах, в том числе в Cut Cross-Entropy. Лестница оптимизаций — тоже ретроспективный разбор, а не воспроизводимый набор бенчмарков: ни исторические ревизии исходников, ни полные параметры запуска не зафиксированы. Результат для прямого прохода не означает такого же сокращения для полного шага обучения.

1. Потоковый подсчёт потерь без полной матрицы логитов

Пусть H — матрица скрытых состояний (N, D), а W — выходная проекция (V, D). Обычная cross-entropy сначала вычисляет все логиты

Z = H Wᵀ

а затем для целевого токена y_i считает

loss_i = logsumexp(Z_i) - Z_i,y_i

Для функции потерь не нужен произвольный доступ ко всей строке логитов. Нужны лишь две построчные статистики: log-sum-exp и логит целевого токена. Это и позволяет перейти к потоковой формулировке.

Ядро сливает каждый тайл словаря в текущий онлайновый log-sum-exp. Пусть пара (m, s) хранит текущий максимум и сумму экспонент, масштабированную этим максимумом. Для нового тайла с максимумом m_t обновление выглядит так:

m' = max(m, m_t)
s' = s · exp(m - m') + Σ_j exp(z_j - m')

После последнего тайла, когда (m, s) — итоговое состояние, logsumexp = m + log(s). Ядро захватывает логит целевого токена, когда его индекс в словаре попадает в текущий тайл. В памяти остаются лишь текущее состояние и итоговые значения по каждому токену. Матрица (N, V) не возникает никогда.

В изучаемой здесь реализации Python формирует один ленивый Metal-запуск (dispatch) на каждый тайл словаря. Ядро JIT-компилируется один раз, а кэшированный конвейер переиспользуется по всей цепочке. При V=151936 и тайле шириной 8192 столбца рабочая форма образует 19 зависимых запусков. Каждый запуск потребляет N-элементные накопители предыдущего и выдаёт следующие. Одно вычисление исполняет и синхронизирует всю готовую цепочку.

Речь идёт только о прямом проходе. Рабочий путь обучения сохраняет log-sum-exp прямого прохода и использует отдельный обратный проход на чистом MLX: он разбит на чанки и заново вычисляет нужные логиты тайл за тайлом. Слитного обратного Metal-ядра там пока нет, поэтому изолированные измерения из этой статьи нельзя читать как результаты для прямого и обратного прохода вместе.

2. Почему одного разбиения на чанки в MLX не хватило

Первый прототип реализовал то же разбиение по словарю обычными операциями MLX. Он был численно корректен, поддерживал плотные и квантованные выходные головы и показал, что потери на большом словаре можно вычислить, не собирая один монолитный тензор логитов.

Он же обнажил контринтуитивное правило вычисления. Вызов mx.eval после каждого чанка, пока MLX трассировал градиент, удерживал трассируемые промежуточные значения вместо того, чтобы освобождать каждый чанк. В раннем эксперименте с прямым и обратным проходом при N=8192 принудительное вычисление на каждом чанке дошло до 55.98 GiB против 9.13 GiB, когда чанки оставались ленивыми, а граф вычислялся один раз в конце. Версия с немедленным (eager) вычислением израсходовала примерно в шесть раз больше памяти.

Когда чанки оставались ленивыми, пиковая память прямого и обратного прохода на измеренных формах снижалась примерно вдвое. Однако при N=4096 это обходилось в 1.31× относительно плотного пути для плотной головы и в 1.71× — для квантованной. Эти разведочные результаты и подтолкнули к написанию слитного ядра. Это не финальные бенчмарки библиотеки, и их нельзя напрямую сравнивать с приведёнными ниже числами только для прямого прохода.

Разбиение на чанки выражало верную математическую декомпозицию, но не могло управлять временем жизни и слиянием каждого промежуточного значения. Собственное Metal-ядро способно держать логиты и редукции отдельного тайла в регистрах, превращая желаемую границу по памяти в свойство самой реализации, а не в надежду на удачное планирование.

Three forward paths to the same per-token lossMaterialized MLXPure-MLX chunkingFused Metal forwardHidden states H (N x D)Output weights W (V x D)Matrix multiplyWrite logits Z (N x V)to device memoryRow log-sum-expand target lookupPer-token loss (N)Hidden states H (N x D)One vocabulary tile of WBuild MLX operationsfor tile logits Z_t (N x C)MLX manages tile intermediatesMerge (m, s, target)Repeat for eachvocabulary tilePer-token loss (N)Hidden states H (N x D)One vocabulary tile of WMetal dispatch computestile logits inside the kernelTile logits remain kernel-localWrite next (m, s, target), all O(N)Chain one dispatch pervocabulary tilePer-token loss and LSE (N)Conceptual dataflow, not a measured allocation-lifetime trace.
Рисунок 1. Три пути вычисляют одни и те же потери по каждому токену, но по-разному обходятся с промежуточным состоянием (подписи на диаграмме — на английском). Это концептуальная схема потоков данных, а не измеренная трасса времени жизни аллокаций. Редактируемый исходник PlantUML опубликован вместе со статьёй.

3. Условия экспериментов

Все приведённые измерения получены на одном Apple M1 Max с 32 GB единой памяти. Разведочная лестница ядер использовала MLX 0.31.2 и macOS 26.5.1. Зафиксированное в репозитории воспроизведение слоя потерь — MLX 0.32.0 и macOS 26.5.2. В обоих случаях входы были в bfloat16 при V=151936 и D=4096.

Каждый замеряемый вариант запускался в отдельном новом процессе. Перед замером он инициализировал и вычислял свои входы, затем прогревал JIT-компилятор и аллокатор. В таблице приведена медиана трёх синхронизированных прогонов.

Перед бенчмарком каждый вариант ядра должен был совпасть с эталоном прямого прохода в fp32. Разведочные наборы покрывали от 10 до 19 случаев на вариант, включая краевые случаи — хвосты и выравнивание. Наибольшее зафиксированное абсолютное расхождение потерь на токен составило 4.8e-6 для входов fp32 и 1.9e-6 для bfloat16. Эти проверки покрывают только прямой проход, но не корректность градиента.

Разведочные артефакты проверки на совпадение (parity) не опубликованы. Финальная библиотека публикует бенчмарк прямого прохода с версиями MLX и ОС, идентификацией пакета и исходника, формой, числом повторов и сырыми замерами времени по часам (wall-clock). Этот артефакт появился раньше поля script_sha в драйвере, поэтому по нему нельзя установить точную ревизию скрипта, породившую строки.

Зафиксированный прогон на MLX 0.32.0 дал такой изолированный результат для прямого прохода:

реализация активная до сброса прирост после сброса общий пик после прогрева медианное время (wall) производительность
Материализованные логиты MLX 5.8585 GiB 2.3184 GiB 8.1768 GiB 1.285595 s 3965.577 G MAC/s
Слитное Metal-ядро 1.2218 GiB 0.0006 GiB 1.2224 GiB 2.104592 s 2422.382 G MAC/s

Замеры времени в одной сессии дают замедление прямого прохода в 1.64×. Со столбцами памяти нужно быть аккуратнее. Бенчмарк прогревает каждую реализацию, сохраняет живыми потери с прогрева, очищает кэш аллокатора и обнуляет счётчик пика. Поэтому в начале окна измерения у материализованного пути активной памяти на 4.6367 GiB больше, чем у слитного.

Этот разрыв в базовом уровне искажает отношение между двумя столбцами прироста, а значение 0.0006 GiB к тому же округлено до четырёх знаков после запятой. Таблица показывает разное поведение памяти, но не даёт чёткого ответа, во сколько раз сокращается память на один вызов. Для такого ответа понадобились бы сырые счётчики байтов и новый прогон, где оба пути стартуют с одинакового базового уровня: в памяти только входы, без живого графа после прогрева.

4. Модель повторного использования загрузок срабатывает лишь раз

Первоначальная Metal-конструкция назначала одну SIMD-группу (группа из 32 потоков, аналог варпа) на несколько строк. Каждая дорожка (lane) проходила по скрытой размерности, накапливала частичные скалярные произведения в fp32 и объединяла дорожки через simd_sum. Главной экспериментальной переменной был регистровый тайл: сколько строк (R) и столбцов словаря (C) одна SIMD-группа обрабатывает вместе.

Тайл R × C с четырёхэлементными векторными загрузками bfloat16 выполняет 4RC операций умножения с накоплением (multiply-accumulate, MAC), запрашивая при этом (R+C) векторов из внутреннего цикла на уровне исходного кода. Мы определили номинальный коэффициент повторного использования загрузок:

load reuse = 4RC / (8(R+C)) MAC per requested byte.

Из коэффициента следовало, что переход от тайла 4 × 1 к 4 × 4 может поднять производительность примерно в 2.5×. При N=8192 измеренная производительность выросла с 294.8 до 814.9 G MAC/s — прирост в 2.76×. На первой паре предсказание совпало и по направлению, и по масштабу. Этот коэффициент — не аппаратная арифметическая интенсивность: он игнорирует записи (stores), работу онлайн-редукции, трафик на уровне кэша, трафик аккумуляторов и состав инструкций. И всё же было заманчиво принять первое совпадение за предсказательную модель.

Следующий тайл оказался медленнее. Переход от 4 × 4 к 4 × 8 поднял номинальное повторное использование загрузок на 33%, но снизил производительность на обеих измеренных длинах последовательности.

В таблице для каждой экспериментальной ступени использованы исторические метки. Нет опубликованной неизменяемой таблицы, которая связывала бы каждую метку с ревизией исходника и полными параметрами запуска. В частности, геометрия тайла по словарю и геометрия запусков (dispatch) для каждой строки не сохранены. Таблица годится для ретроспективного сравнения, но не для независимого повторного прогона лестницы.

вариант регистровый тайл номинальное повторное использование загрузок производительность при N=8192 потолок потоков при компиляции
v1c 4 × 1 0.40 MAC/B 294.8 G MAC/s не записан
v1d 4 × 4 1.00 MAC/B 814.9 G MAC/s 512
v1f 4 × 8 1.33 MAC/B 752.3 G MAC/s 384
v1e 8 × 8 2.00 MAC/B безопасный запуск на этой форме невозможен 384

У тайла 4 × 8 было на 33% больше номинального повторного использования загрузок, чем у 4 × 4, и всё же он был на 8% медленнее при N=8192 и примерно на 20% медленнее при N=2048. Версия 8 × 8 просела ещё сильнее: в сравнении на одной и той же форме и одном и том же тайле в 4096 столбцов при N=2048 она выдала 182.1 G MAC/s против 635.0 у 4 × 8.

5. Четыре контрольных эксперимента сужают круг причин

Несколько объяснений выглядели правдоподобно. Шаг строки (stride) был ровно 8192 байта. Это совпадает с размером кэша данных L1 у GPU M1 по полученным обратной разработкой заметкам об архитектуре metal-benchmarks, так что дополнительные параллельные потоки данных могли вызывать конфликты по множествам кэша. Часть вариантов изначально использовала встроенную копию измерительной обвязки. Эксперимент 8 × 8 к тому же изменил и форму тайла, и стиль исходного кода. Его первое сравнение на малой форме не дало обоим вариантам одинаковой возможности загрузить GPU.

Сузить поле помогли четыре контрольных эксперимента:

  1. Смена D с 4096 на 4160 изменила шаг строки с 8192 на 8320 байт. Отношение производительности 4 × 8/4 × 4 осталось 0.80 при N=2048. Это исключило гипотезу о конфликтах наборов кэша из-за шага, равного степени двойки, как причину регрессии 4 × 8.
  2. Прогон обоих вариантов подряд через общий скрипт воспроизвёл прежние показатели с точностью до 1%, что исключило скопированную обвязку как причину.
  3. Переписывание ядра 4 × 4 в стиле «массив и разворачивание цикла», который использовали тайлы покрупнее, воспроизвело реализацию с явными скалярами. Идиома записи кода была ни при чём.
  4. Сравнение 8 × 8 и 4 × 8 при одинаковых N=2048 и размере тайла по словарю оставило остаточное замедление в 3.5×. Это убрало разницу в числе запусков и уменьшило проблему насыщения на малой форме, но у вариантов всё ещё была разная геометрия блоков строк.

Мы также скомпилировали сгенерированный Metal-исходник средствами фреймворка Metal и изучили MTLComputePipelineState.maxTotalThreadsPerThreadgroup. Максимум устройства — 1024 потока. Ядро 4 × 4 скомпилировалось с потолком 512, а 4 × 8 и 8 × 8 — оба на 384. Этот предел допустимости конвейера коррелирует с более тяжёлым расходом ресурсов в скомпилированном ядре. Он не измеряет ни достигнутую занятость (occupancy), ни резидентные SIMD-группы, ни регистры. Эксперимент также не показал, что сниженный потолок менял резидентность (residency) для реального запуска.

Контрольные эксперименты подтверждают, что регрессии реальны, и связывают их с более тяжёлым расходом ресурсов в скомпилированном ядре. Причину они не называют. Для 4 × 8 один из кандидатов — сниженная занятость. Для 8 × 8 оценённый объём живого состояния превысил 128 регистров по 32 бита (GPR) — предел, известный из обратной разработки в заметках об Apple GPU Дугалла Джонсона. После согласованных контролей осталось большое замедление. Вытеснение регистров в память (register spill) — кандидат, но ни ISA компилятора, ни статистики вытеснений, ни счётчика трафика памяти, ни измерения достигнутой занятости мы не собрали. Причинная связь ни для одного из механизмов здесь не доказана.

6. Матричные тайлы уменьшают скалярное состояние

Следующая конструкция использовала операции simdgroup_matrix, чтобы повысить повторное использование, не наращивая обычное скалярное состояние таким же образом. Её ранние ступени были медленнее лучшего ядра на регистровых массивах. Более крупные матричные тайлы со временем закрыли этот разрыв.

ступень конструкция производительность при N=8192 замедление к материализованному пути MLX
v2a один матричный тайл 8 × 8 487.2 G MAC/s 8.1×
v2c тайлы 2 × 2 (16 × 16) 1233.9 G MAC/s 3.2×
v2d тайлы 2 × 4 (16 × 32) 1579.3 G MAC/s 2.5×
v2e тайлы 4 × 4 (32 × 32) 2423.7 G MAC/s 1.63×
v2f тайлы 4 × 8 (32 × 64) 1403.2 G MAC/s 2.8×

Производительность росла вплоть до тайла 32 × 32, который использовал примерно 32 fp32-элемента аккумулятора на дорожку. Удвоение одной из сторон тайла затем срезало производительность с 2423.7 до 1403.2 G MAC/s. В обоих семействах конструкций рост повторного использования данных помогал лишь до тех пор, пока дополнительное состояние на дорожку не делало скомпилированное ядро заметно тяжелее.

Компактная модель наблюдаемого поведения такова:

производительность может ограничиваться как доставкой данных, так и расходом ресурсов в скомпилированном ядре.

Это качественная эвристика проектирования, а не подогнанная модель и не универсальный закон для GPU Apple. Эксперимент не измерял ни трафик памяти, ни потолки пропускной способности памяти, ни число регистров, ни вытеснения, ни достигнутую занятость, поэтому он не может разделить вклад пропускной способности памяти, темпа инструкций загрузки (load), поведения кэша и сокрытия задержек. Вывод скромнее: более высокое номинальное повторное использование загрузок всё равно может проиграть, если делает ядро тяжелее каким-то другим, неизмеренным образом.

7. Экономия на прямом проходе не предсказывает экономию на всём шаге

Слитное ядро по своей конструкции убирает выделение логитов из слоя потерь прямого прохода. Публичный прогон по памяти согласуется с большим сокращением, но его неодинаковые базовые уровни после прогрева не позволяют получить чистое отношение. В полном шаге обучения остаются трансформер, оптимизатор, веса модели и обратный проход внимания. Снятие одного выделения не гарантирует пропорционального роста максимального контекста.

Измерения из релиза 0.1.0 проекта показывают, почему важен охват. Запись о релизе сообщает: при одном и том же штатном (stock) внимании во всех вариантах сравнения одна лишь слитная cross-entropy не дала заметного роста максимального обучаемого контекста. Пик на длинном контексте она относит к обратному проходу штатного внимания. Поскольку сырой исторический прогон не опубликован, это остаётся результатом, известным только из отчёта проекта, а не независимо воспроизводимым сравнением.

Таким образом, ядро потерь убирает лишний расход памяти в слое потерь, но не решает проблему памяти при обучении в целом. Само по себе оно не увеличивает максимальный контекст.

Скорость тоже зависит от охвата. Изолированный слитный прямой проход был в 1.64× медленнее тонко настроенного материализованного пути MLX, но слой потерь — лишь часть шага обучения. Более поздние сквозные измерения нашли куда меньший штраф на уровне шага. Этот штраф зависит от модели, длины последовательности, реализации внимания и обратного пути. Из одной изолированной таблицы его вывести нельзя.

Публичный bench_train_step.py документирует более поздний метод на уровне шага. Вывод про контекст только по слою потерь приведён среди известных ограничений 0.1.0 в changelog. Текущий northstar_context_sweep.py меняет и внимание, и потери. Поэтому он документирует более позднее сравнение продукта целиком, а не этот результат только по слою потерь.

8. Уроки несработавших моделей

Для будущей работы над ядрами MLX:

9. Ограничения

Это исследование использовало один M1 Max, одну основную рабочую форму, входы в bfloat16 и зафиксированные версии MLX. Абсолютные показатели зависят от машины, ОС, релиза MLX, теплового состояния и JIT-компилятора. Медианы по трём прогонам улавливают крупные различия конструкций, а не тонкую дисперсию.

Эксперименты не подтвердили объяснения через занятость и вытеснение регистров доказательствами на уровне ISA или счётчиков. Публичный протокол по памяти к тому же использует неодинаковые базовые уровни после прогрева, поэтому он не может обосновать чистое отношение памяти на один вызов. Эта статья сообщает об этом ограничении, а не перезапускает эксперимент. Она также охватывает только ядро прямого прохода. Рабочий обратный путь и варианты с квантованной головой — вне её охвата.

10. Воспроизводимость и источники

Финальный бенчмарк слоя потерь порождается скриптом bench_loss_layer.py. Зафиксированный в коммите JSON-артефакт хранит идентификаторы его условий и все три замера времени по часам (wall-clock). У разведочной лестницы нет аналогичного публичного набора, поэтому её результаты известны только из отчёта проекта и их нельзя воспроизвести независимо. Публичный JSON указывает измеренный исходник пакета, но появился раньше поля script_sha в драйвере. Поэтому точная историческая ревизия драйвера неизвестна.

Cut Cross-Entropy развивает более общий подход линейной cross-entropy без материализации логитов. Liger-Kernel предоставляет слитные ядра обучения для других стеков ускорителей.

11. Заключение

Слитный прямой проход убирает матрицу логитов (N, V), но экономит память ценой времени. На измеренной форме M1 Max он работал в 1.64× медленнее материализованного MLX. Публичный результат по памяти указывает на большое сокращение, но его неодинаковые базовые уровни после прогрева не дают надёжного отношения на один вызов.

Лестница оптимизаций важна по другой причине. Модель повторного использования загрузок верно предсказала улучшение 4 × 4 и ошиблась на 4 × 8; матричный тайл 32 × 32 позже вернул лучшую наблюдавшуюся производительность, прежде чем более крупный тайл снова дал регрессию. Без счётчиков регистров или занятости причина остаётся открытой. Обоснованный результат более узкий, но полезнее: повторное использование данных помогает лишь до тех пор, пока расход ресурсов в скомпилированном ядре под контролем, а одно лишь повторное использование на уровне исходного кода не может предсказать эту границу.


Подготовлено 2026-07-14. Последнее обновление 2026-07-19. Denis Ineshin.

Buy me a coffee