Ссылка: Технический отчет Kimi Linear — https://github.com/MoonshotAI/Kimi-Linear/blob/master/tech_report.pdf
Код:
-
tokamax/_src/ops/experimental/kda/base.py— публичный аргумент KDA и контракт рекуррентного состояния -
tokamax/_src/ops/experimental/kda/reference.py— реализация рекуррентной эталонной модели на чистом JAX -
tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu.py— Пользовательский адаптер VJP для Tokamax -
tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_kernel.py— общедоступные точки входа в ядро нижнего уровня -
tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_types.py— типизированный контрактKdaResiduals -
tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_fwd_kernel.py— построение матриц прямых остатков и сохраненных состояний -
tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_bwd_kernel.py— обратный оркестратор и ядра Pallas -
tokamax/_src/ops/experimental/kda/common.py— вспомогательная функция для работы с суммой по общим вентилям, рекуррентным соотношением состояний и мини-пакетами. -
tokamax/_src/ops/experimental/kda/cp_utils.py— контекстно-параллельное слияние градиентов -
tokamax/_src/ops/experimental/kda/utils.py— вспомогательные функции для выравнивания, расвыравнивания и обратной L2-нормализации
1. Цель и поток данных обратного прохода
KDA обрабатывает последовательность блоками, состоящими из временных шагов BT , при этом BT=64 по умолчанию.
Для каждого фрагмента прямой выходной сигнал представляет собой сумму двух частей:
\mathbf{o}
= \underbrace{s\,(\mathbf{q}\odot 2^{\mathbf{g} })\,\mathbf{h} }_{\text{inter: reads historical state} }
+ \underbrace{\mathrm{tril}(\mathbf{A}_{qk})\,\mathbf{v}_{\text{new} } }_{\text{intra: reads earlier tokens within the chunk} }
Интуитивно понятно, что результат запроса поступает по двум путям:
- inter path: считывает информацию из исторического состояния
h, которое уже существует при входе в этот фрагмент; - Внутрипутевой поток: считывает информацию из обновленных результатов предыдущих токенов в этом блоке.
Ввод Tokamax — это PallasMosaicTpuKimiDeltaAttentionVjp._fwd(...) , который передает типизированные прямые остатки и выходные котангенсы в chunk_kda_bwd_custom(...) . Оркестратор получает выходной градиент do и необязательный градиент конечного состояния dht и выдает:
-
dq, dk, dv: градиенты относительно запроса/ключа/значения; -
db: градиент относительноbeta; -
dg: градиент относительно входного сигнала публичногоgate. Обратный кумулятивный сумматор объединенного ядра сначала вычисляет градиент активированного вентиля для каждого токена, а затем кумулятивный сумматор; когдаuse_gate_in_kernel=True,kda_gate_bwd(...)отображает этот градиент дальше обратно на исходный публичный вентиль; -
dh0: градиент относительно начального состояния, возвращается только в том случае, если вызывающая сторона предоставляетinitial_state; -
dA, dbias: градиенты относительноa_logи необязательныйdelta_time_bias, когда активация вентиля вычисляется внутри ядра.
dht — это необязательный входной параметр, представляющий градиент конечного состояния. Он отличен от нуля только тогда, когда конечное состояние продолжает использоваться последующим слоем потерь; когда обычный слой внимания обрабатывает только выход o , dht может быть None и рассматривается как ноль.
1.1 Перенос сохраненных значений вперед и обратный пересчет границы
Обратный проход не выполняет все вычисления прямого прохода с нуля. Текущая реализация делит необходимые для обратного прохода величины на три категории.
Сохранены прямые остаточные значения:
-
Aqk: матрица внимания к ключам запроса внутри блока, форма[H, B, T, BT]; -
Akk: обратная матрица WY, т.е.A_kk^{-1}, форма[H, B, T, BT]; -
h: скрытое состояние в начале каждого фрагмента, форма[H, B, NT, K, V]; при перенаправлении оно сохраняется вKdaResidualsеслиrematerialize_for_backward=False.
Обратный перерасчет:
-
w: WY эффективный вес клавиши / вес клавиши стирания; -
qg: запрос с ограниченным доступом; -
kg: ключ с воротами; -
v_new: значение, скорректированное по стандарту WY.
Дополнительный перерасчет:
- Когда
use_gate_in_kernel=True, оба пути обратной подготовки в настоящее время пересчитывают пост-кумум-вентиль изg_org,a_logи необязательногоdelta_time_biasчерезkda_gate_chunk_cumsum, включая путь saved-h, где прямой остаток в настоящее время также содержитg_cumsum. - Если
rematerialize_for_backward=True, обратный путь не следует по сохраненному путиh, а вместо этого повторно запускает рекуррентное выражение прямого состояния для восстановленияhиv_new.
Это ручная рематериализация внутри пользовательского VJP Mosaic, а не вызов jax.checkpoint или jax.remat .
1.2 Представление WY и семантика вентилей
Обновление состояния внутриблочного правила дельта-изменения по своей сути является последовательным: каждый токен управляет своей силой записи с помощью beta , а последующий токен зависит от состояния, обновленного предыдущим токеном.
В представлении WY это последовательное обновление кодируется в нижнетреугольную матрицу A_kk , а для преобразования внутриблоковых вычислений в параллельные матричные умножения используется обратная матрица A = A_kk^{-1} сохраненная в прямом порядке:
\mathbf{u}=\mathbf{A}(\mathbf{v}\odot\boldsymbol\beta)
\mathbf{w}=\mathbf{A}(\mathbf{k}\odot\boldsymbol\beta\odot 2^{\mathbf{g} })
\mathbf{v}_{\text{new} }=\mathbf{u}-\mathbf{w}\,\mathbf{h}
Akk — это сохраненная обратная матрица (I+L)^{-1} . Строго нижнетреугольная матрица L уже зависит от построчного значения beta , и beta также явно умножается на правые части значений и ключей с вентилями, используемые для построения u и w . Следовательно, обратный путь должен учитывать оба пути, по которым beta влияет на результат.
В приведенных ниже уравнениях g обозначает кумулятивный вентиль, локальный для каждого блока, в логарифмическом пространстве log2. Это внутренняя величина, а не публичный аргумент вентиля. Поэтому все распады записываются как 2^(...) и вычисляются в ядре с помощью exp2 :
-
qg = q * 2^g: запрос считывает историческое состояние, изменяющееся от начала фрагмента до текущей позиции; -
kg = k * 2^(g_C - g): после записи ключа он должен исчезнуть до конца фрагмента, прежде чем сможет перейти в состояние следующего фрагмента; -
2^g_C: распад всего состояния, проходящего через один фрагмент, гдеg_C— кумулятивный вентиль в конце фрагмента.
Соответствующая межблоковая повторяемость выглядит следующим образом:
\mathbf{h}^{[t+1]}
=2^{\mathbf{g}_C}\odot\mathbf{h}^{[t]}
+\mathrm{kg}^\top\mathbf{v}_{\text{new} }
1.3 API и ссылки на переменные
Точка входа: tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu.py::PallasMosaicTpuKimiDeltaAttentionVjp._fwd(...) , за которым следует tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_kernel.py::chunk_kda_bwd_custom(...) .
Все тензоры основного пути используют структуру [H, B, T, ...] с направлением «сверху вниз». Публичный API KDA и бэкенд Pallas уже используют эту структуру; на границе пользовательского VJP нет транспонирования входных данных. Обратный путь использует выровненные и, при необходимости, L2-нормализованные копии, хранящиеся в KdaResiduals , а не воспроизведенные исходные тензоры, предоставляемые универсальным контрактом VJP Tokamax.
Входные данные и значения, сохраненные в дальнейшем:
| Переменная | Форма | Значение |
|---|---|---|
публичный query, key / остаточный q, k | [H, B, T, K] | запрос / ключ |
общественная value / остаточная v | [H, B, T, V] | ценить |
beta | [H, B, T] | дельта-правило — сила записи на токен |
общественные gate | [H, B, T, K] | входной сигнал для каждого токена перед вычислением локальной суммы блока |
g_cumsum | [H, B, T, K] / None | Внутренний пост-кумумный вентиль в логарифмическом пространстве по основанию 2; пересчитывается при пропуске. |
g_org | [H, B, T, K] / None | Сохранение необработанных данных при выполнении активации в ядре |
a_log | [H] / None | параметр активации затвора |
delta_time_bias | [H*K] / None | необязательное смещение активации затвора |
Aqk | [H, B, T, BT] | матрица внимания внутри блока с сохранением данных при пересылке |
Akk | [H, B, T, BT] | обратная матрица WY, сохраненная в прямом направлении |
h | [H, B, NT, K, V] | скрытое состояние в начале каждого фрагмента, сохраняется на быстром пути. |
do | [H, B, T, V] | выходной градиент |
initial_state | [B, N, H, K, V] / None | исходное состояние в рамках публичного контракта KDA |
dht | [B, N, H, K, V] / None | котангенс конечного состояния в рамках публичного контракта KDA |
Обратные промежуточные величины:
| Переменная | Форма | Значение |
|---|---|---|
u | [H, B, T, V] | Эффективное значение WY существует лишь временно внутри ядра пересчета. |
w | [H, B, T, K] | Эффективный ключ WY / вес стирания |
qg, kg | [H, B, T, K] | ограниченный запрос / ключ |
v_new | [H, B, T, V] | значение после WY и корректировки на основе исторического состояния |
dAqk | [H, B, T, BT] | градиент матрицы внутриобъектного внимания, полученный из ядра dAv. |
Результаты:
| Переменная | Форма | Значение |
|---|---|---|
dq, dk, dg | [H, B, T, K] | запрос / ключ / публичный градиент входа вентиля |
dv | [H, B, T, V] | градиент значений |
db | [H, B, T] | beta градиент |
dh0 | [B, N, H, K, V] / None | Градиент начального состояния, соответствующий форме общедоступных входных данных. |
dA, dbias | [H] / [H*K] / None | градиенты a_log / delta_time_bias |
Статические ограничения:
- В конфигурации Mosaic, предоставленной производителем, указано
BT=chunk_size=64; - Подготовленный
Tдолжен делиться наBT; входные данные переменной длины выравниваются перед построением остатка; - Основной путь использует логический вентиль log2, т.е.
use_exp2=True; - Выполнение без CP поддерживает контракт
K<=256, предоставляемый бэкэндом, включаяK/V, не выровненные по 128; в настоящее время CP требует, чтобы иK, иVбыли кратны 128.
Обратный оркестратор может нормализовать внутренний четырехмерный котангенс конечного состояния, вставив ось N=1 , но это путь обеспечения совместимости реализации. Публичный контракт Tokamax KDA принимает и возвращает рекуррентные состояния в пятимерной форме [B, N, H, K, V] . Для выполнения Pallas фиксированной длины требуется N=1 .
2. Цепочка вызовов на основе кода и структура ядра
Текущая цепочка вызовов Tokamax выглядит следующим образом:
PallasMosaicTpuKimiDeltaAttentionVjp._fwd
└─ chunk_kda_bwd_custom
├─ unpack KdaResiduals
├─ align do with the retained cu_seqlens/aligned_cu_seqlens
├─ restore derived CP metadata retained by forward
├─ Stage 0: recover forward intermediate quantities
│ ├─ if use_gate_in_kernel: kda_gate_chunk_cumsum(...)
│ ├─ saved-state path rematerialize_for_backward=False:
│ │ └─ fused_recompute_w_u_vnew_from_h_pallas(...)
│ │ → w, qg, kg, v_new
│ └─ low-memory path rematerialize_for_backward=True:
│ ├─ _recompute_w_u_fwd(...)
│ │ → w, u, qg, kg
│ └─ chunk_gated_delta_rule_fwd_h(...)
│ → h, v_new
├─ Stage 1: chunk_kda_bwd_dAv_kernel(...)
│ → dAqk, dv
├─ optional context parallel:
│ ├─ chunk_gated_delta_rule_bwd_dhu_pre_process(...)
│ │ → dS_ext, dM
│ ├─ all_gather_into_tensor(...)
│ ├─ _merge_dht(...)
│ └─ construct the dht used by the fusion kernel
├─ Stage 2-5: _fused_dhu_wy_intra_cumsum_pallas_jit(...)
│ → dq, dk, dv, db, dg, dh0
├─ optional gate backward:
│ └─ kda_gate_bwd(...)
│ → dg, dA, dbias
├─ optional l2norm_bwd(...) for dq/dk
├─ _unalign_output(...) for varlen dq/dk/dv/dg/db
└─ restore dh0 shape and input dtypes
Математически обратный проход можно разложить на шесть этапов:
- пересчитать промежуточные величины;
- вычислить градиент внутричлена
Aqk @ v_new; - повторяйте градиент состояния
dhв обратном направлении вдоль оси фрагмента; - Обратное распространение представления WY;
- распространить обратное распространение внимания внутри блока;
- выполнить обратное суммирование по градиенту вентиля.
Основной путь saved- h сопоставляет эти шесть этапов с тремя ядрами Палласа:
recompute fusion → w, qg, kg, v_new
dAv → dAqk, dv
dhu/WY/intra/cumsum → dq, dk, dv, db, dg, dh0
Выбранный параметром rematerialize_for_backward=True путь с низким потреблением памяти дополнительно вызывает _recompute_w_u_fwd и chunk_gated_delta_rule_fwd_h для повторного выполнения рекуррентного алгоритма прямого состояния. Подготовка контекста, обратное вентилирование, обратная L2-нормализация и выравнивание Варлена являются необязательными операциями по отношению к основным ядрам.
2.1 Перерасчет ядра слияния
Точка входа: fused_recompute_w_u_vnew_from_h_pallas(...) .
Это ядро считывает сохраненные значения h и Akk , сохраненные в дальнейшем, и параллельно пересчитывает следующие значения для каждого блока:
u = Akk @ (v * beta)
w = Akk @ (k * beta * 2^g)
v_new = u - w @ h
qg = q * 2^g
kg = k * 2^(g_C - g)
u существует только внутри ядра и не записывается обратно в HBM после вычисления v_new . Поскольку h уже сохранено, v_new больше не зависит от повторного возникновения состояний между блоками, поэтому каждый блок может быть запланирован независимо.
Путь с меньшим объемом памяти используется, когда h не сохраняется. Сначала он восстанавливает w, u, qg, kg с помощью _recompute_w_u_fwd(...) , затем вызывает chunk_gated_delta_rule_fwd_h(...) , чтобы повторно запустить рекуррентное выражение для состояний и получить h, v_new . Этот путь восстанавливает последовательную зависимость между блоками и обменивает вычислительные ресурсы на меньший объем остаточной памяти для дальнейшего выполнения.
2.2 dAv Ядро
Точка входа: chunk_kda_bwd_dAv_kernel(...) .
Данное ядро обрабатывает только внутриядерный вывод:
\mathrm{tril}(\mathbf{A}_{qk})\,\mathbf{v}_{\text{new} }
Соответствующее обратное распространение ошибки выглядит следующим образом:
d\mathbf{A}_{qk}
=s\,\mathrm{tril}(d\mathbf{o}\,\mathbf{v}_{\text{new} }^\top)
d\mathbf{v}_{\text{new} }
=\mathrm{tril}(\mathbf{A}_{qk})^\top d\mathbf{o}
В коде scale умножается только один раз, при генерации dAqk . Последующее ядро слияния напрямую использует уже масштабированный dAqk , избегая повторного масштабирования.
Сохраненный Aqk уже содержит scale , поэтому ветвь value-gradient использует Aqk^T @ do без дополнительного множителя. Тензор с именем dAqk масштабируется перед объединенным обратным преобразованием intra, поскольку это ядро напрямую различает базовое немасштабированное отношение запрос-ключ. Хотя chunk_kda_bwd_dAv_kernel сохраняет q и k в своей публичной сигнатуре, текущий лаунчер и ядро используют только параметры v , Aqk , do , scale и tiling.
2.3 dhu/WY/intra/cumsum Fusion Kernel
Точка входа: _fused_dhu_wy_intra_cumsum_pallas_jit(...) .
Это ядро выполняется в обратном порядке блоков и объединяет четыре типа работы:
Функция обратной рекуррентности
dhuподдерживает градиент состояния между блокамиdh, накапливая вклады от выходного пути, пути дельта-обновления и пути затухания блока в один и тот же временный массив VMEM.Обратное распространение ошибки WY происходит от
v_new = u - w @ h,u = Akk @ (v * beta),w = Akk @ (...)доq, k, v, beta, g, и позволяет получить вклад градиента вAkk.Функция intra backward потребляет
dAqkполученный ядром dAv, продолжает обратное распространение градиента матрицы внимания внутри блока кq, k, beta, gи накапливает результаты WY.Обратная кумулятивная сумма. Прямой проход использует локальную кумулятивную сумму блока, поэтому для обратного прохода требуется обратная кумулятивная сумма. Текущая реализация хранит это в рамках одного ядра и больше не запускает отдельное ядро для вычисления кумулятивной суммы.
Ось блоков использует arbitrary семантику сетки для передачи состояния обратной рекуррентности; оси заголовка и пакета являются параллельными измерениями.
2.4 Дополнительный контекстный параллельный режим
Когда включена параллельная обработка контекста, обратный проход вставляет слияние градиентов состояний по рангам между ядрами dAv и fusion.
Последовательность действий следующая:
-
chunk_gated_delta_rule_bwd_dhu_pre_process(...)сканирует только первый реальный локальный сегмент, поскольку именно этот сегмент может получать состояние от вышестоящего уровня, и выдает следующий результат:-
dS_ext: внешний вклад этого ранга в градиент входного состояния; -
dM: произведение матриц обратных переходов этого ранга.
-
- Упакуйте
dS_extиdMвдоль последнего измерения и обменяйте их однимall_gather_into_tensor(...). - Функция
_merge_dht(...)объединяет вклады нижестоящих рангов в порядке от наиболее удаленного к наиболее удаленному, используя метаданныеpost_num_ranksиis_last_rank, сохраняемые функцией forward. - Создайте
dht[B,N,H,K,V]и поместите градиент объединенного состояния каждого элемента пакета в его последний локальный сегмент перед входом в объединенное ядро обратного преобразования.
Обратное направление CP идет от нисходящего ранга к восходящему, определяемому выровненными segment_ids и метаданными ранга, восстановленными в ContextParallelMetadata . B>1 обрабатывается независимо для каждого элемента пакета. Контракт CP не принимает внешнее initial_state и не возвращает градиент начального состояния.
2.5 Дополнительные ворота, открывающиеся назад
Когда use_gate_in_kernel=True , chunk_kda_bwd_custom(...) вызывает kda_gate_bwd(...) после обратного хода ядра.
На этом этапе ядро слияния уже применило локальную для блока обратную кумулятивную сумму, поэтому его dg представляет собой градиент относительно активированного вентиля для каждого токена до применения кумулятивной суммы. Затем kda_gate_bwd(...) дифференцирует активацию и отображает этот градиент на исходный публичный gate a_log и необязательный delta_time_bias , возвращая соответственно dg , dA и dbias .
3. Адаптация TPU
Цель этих решений по реализации состоит в том, чтобы обратный путь максимально точно соответствовал модели выполнения MXU/VPU/HBM процессора TPU.
3.1 log2 Gate и exp2
Суммарное значение, потребляемое ядрами внутрисостоятельной и рекуррентной рекуррентности, представляется в логарифмическом масштабе по основанию 2, поэтому убывания равномерно записываются как 2^g . Входной сигнал открытого вентиля преобразуется в это представление на этапе прямого вентиля/суммирования; его не следует путать с внутренним суммарным значением.
Это дает два преимущества:
- Внутри- и межчастичные распады можно объединить путем сложения и вычитания в показателе степени;
- Ядро использует аппаратный
exp2, что позволяет избежать многократной смены базового значенияexp.
3.2 fp32 Накопление и bf16 Хранение
В качестве входных данных обычно используется bf16, но для умножения матриц используется fp32 accumulate.
Внутри основного ядра слияния входные данные fp32 выбирают jax.lax.Precision.HIGHEST , в то время как входные данные bf16 используют точность по умолчанию с накоплением fp32. Обратная сумма точек явно использует jax.lax.Precision.HIGHEST для обоих типов входных данных. Эти выборы важны для dg , где кумулятивные вклады вентилей могут практически компенсировать друг друга.
3.3 VMEM Scratch сохраняет промежуточные продукты кросс-стадионного режима
Ядро слияния помещает обратное состояние dh в временный файл VMEM размером [MB, K, V] и обновляет его в обратном порядке фрагментов.
В то же время, dv_new , промежуточные градиенты WY, промежуточные градиенты intra и локальные результаты, необходимые для обратного вычисления cumsum, остаются максимально циркулирующими внутри VMEM. Это позволяет избежать многократной записи и чтения из HBM между логическими этапами dhu, WY, intra и cumsum.
3.4 Управление мини-пакетами, гранулярность DMA
Программа Pallas обрабатывает по одному фрагменту данных MB head/chunk tiles) за раз.
MB выбирается из предполагаемого объема памяти VMEM, но точная реализация различается в зависимости от ядра. Функция chunk_kda_bwd_dAv_kernel использует общую вспомогательную estimate_mini_batch(...) . Ядро пересчета saved h , основное ядро слияния и предварительная обработка CP в настоящее время используют локальные эвристики бюджета VMEM с ограничениями, специфичными для ядра, и корректировками делимости/выравнивания.
- Слишком малый размер приводит к недостаточной детализации DMA и низкой эффективности использования полосы пропускания HBM;
- Слишком большой размер увеличивает нагрузку на VMEM и регистры, а также накладные расходы на статическую развертку.
Таким образом, каждый пусковой модуль оценивает свою площадь тайлов и выбирает MB с учетом требований к сетке, делимости и выравниванию по малым размерам TPU. Пусковой модуль, выполняющий пересчет сохраненных h может дополнять свой сглаженный счетчик чанков, если подходящий делитель недоступен; пусковые модули основного слияния и CP уменьшают MB до тех пор, пока не разделят количество голов.
3.5 Перемещение данных BlockSpec и явный CP DMA
Три основных ядра с сохраненным h используют сетки Pallas и объекты BlockSpec для описания тайлов HBM. Их средства запуска не содержат явного написанного вручную асинхронного цикла копирования.
Предварительная обработка CP отличается: _chunk_gated_delta_rule_bwd_dhu_pre_process_kernel — это активная реализация CP, которая явно использует pltpu.make_async_copy , семафоры DMA и входные данные с двойной буферизацией VMEM. Это однопрограммное обратное сканирование групп голов и блоков, перекрывающее следующую передачу входных данных текущей матричной работой и асинхронно записывающее сводку dS_ext/dM для каждой группы голов.
3.6 Принципы укладки плитки и облицовки
В технологии TPU-тайлинга определенные размеры в конце блока выравниваются по границам аппаратного тайла. Хотя заполненные элементы семантически равны нулю, они все равно потребляют реальную пропускную способность на уровне HBM и DMA.
Текущая реализация основана на практическом принципе: более сложная упакованная структура вводится только тогда, когда объем трафика HBM, который можно исключить путем изменения структуры, превышает дополнительные затраты на изменение формы/сбор/копирование.
Поэтому:
- Обратная поддержка non-CP позволяет обрабатывать невыровненные случаи
K/Vбез наложения ограничения полосы CP, в то время как предварительная обработка CP требует, чтобыKиVбыли кратны 128; -
[BT, BT]небольшие матрицы, такие какAqkиAkkпринимают фиксированное заполнение, обеспечиваемое аппаратным выравниванием; - Скалярные дополнительные измерения, такие как
beta/dbиспользуют явные одномерные или двухмерные макеты, выбираемые каждым лаунчером.
Другими словами, сам по себе отступ не обязательно является чем-то плохим, что нужно обязательно устранять. Пока затраты на индексацию и преобразование макета, возникающие из-за устранения отступов, выше, сохранение простого макета на самом деле быстрее.
3.7 Нормализация субблоков внутри обратного направления
Внутригрупповой обратный процесс включает в себя члены затухания, такие как 2^(g_r - g_j) . Поскольку g является кумулятивной величиной, прямое вычитание и последующее возведение в степень может вызвать проблемы с числовым диапазоном.
В данной реализации фрагмент разбивается на более мелкие подблоки, и в каждом подблоке выбирается опорный вентиль для нормализации. Таким образом, показатель степени разделяется на две части относительно опорной точки, что позволяет удерживать промежуточные значения в более контролируемом диапазоне.
Та же самая структура подблоков также позволяет удобно организовывать диагональные и недиагональные блоки в фиксированное количество пакетных матричных умножений, избегая написания квадратичного цикла по подблокам.
3.8 Фиксированные формы, поддерживающие varlen
Ядра TPU требуют статических форм. Последовательности переменной длины не запускают отдельное неровное ядро, но по-прежнему используют фиксированные тензоры [B, T] и фиксированную сетку (H // MB, B, NT) .
Адаптер вычисляет cu_seqlens , выравнивает каждый логический сегмент по BT и сохраняет как исходные, так и выровненные метаданные в KdaResiduals . Сначала do обратное выравнивание, затем используются сохраненные выровненные segment_ids и сопоставление фрагментов:
- Фрагменты с идентификатором сегмента 0 являются заполнением;
- Последний фрагмент каждого вещественного сегмента инициализирует обратную рекуррентную формулу с помощью
dht; - Первый фрагмент каждого реального сегмента выдает
dh0.
Для наглядности рассмотрим B=1, T=12, BT=4 , содержащие две последовательности длиной 8 и 4. В поставляемом адаптере по-прежнему требуется BT=64 ; меньшее значение лишь делает пример метаданных более компактным.
token segment_ids = [1,1,1,1, 1,1,1,1, 2,2,2,2]
chunk_seg_ids = [ 1, 1, 2 ]
chunk0 chunk1 chunk2
Для выполнения с фиксированной длиной _fused_dhu_wy_intra_cumsum_pallas_jit синтезирует тензор идентификаторов сегментов [B,T] , состоящий из одних единиц, поэтому каждый элемент пакета рассматривается как один сегмент. После обратного преобразования ядра градиенты токенов переменной длины не выравниваются относительно исходной структуры [B,T_original] .
4. Анализ производительности термоядерного синтеза
Измерения в этом разделе сохранены как исторические данные об оптимизации. Они предшествуют текущему адаптеру Tokamax, централизованному выравниванию/остаточному контракту, текущей передаче метаданных CP и некоторым изменениям в именовании на уровне запуска. Перед использованием в качестве порогового значения или базового уровня регрессии их необходимо повторно выполнить для текущего коммита.
Историческая конфигурация измерений: TPU v6e, bf16, varlen seq_lens=[1800,1500,2000,1200,800] , H=16, B=1, T=8192, K=V=128, BT=64 . После разбиения на блоки, начальное измерение составляет 16 × 133 .
4.1 Почему Fusion является основным направлением оптимизации
Большинство этапов обратного преобразования KDA имеют низкую арифметическую интенсивность и, при разделении на отдельные этапы, легко ограничиваются полосой пропускания HBM.
Если каждый математический этап преобразовать в независимое ядро, то большое количество промежуточных тензоров будет передаваться туда и обратно в HBM, например:
-
w, qg, kg, v_new; -
dAqk, dv; -
dh, dv_new; - градиенты WY и внутрипространственные промежуточные градиенты;
- Входные и выходные данные обратной кумулятивной суммы.
В текущей реализации эти этапы сжаты в три основных ядра Pallas. Главное преимущество заключается не в сокращении математических операций, а в уменьшении операций чтения и записи на границах этапов.
4.2 Исторические измерения основного пути с тремя ядрами
Основной маршрут состоит из трех частей:
Измерение: TPU v6e , bf16, varlen seq_lens=[1800,1500,2000,1200,800] (N=5 сегментов), H=16, B=1, T=8192, K=V=128, BT=64; после разбиения на блоки начальное измерение = 2128 = 16×133 (NT=133 = 128 базовых блоков + 5 сегментных дополнений). Пропускная способность HBM принята равной 1,6 ТБ/с.
| Ядро | Время выполнения (мкс) | HBM (MB) | Нижняя граница полосы пропускания (мкс) | использование БВт |
|---|---|---|---|---|
fused_recompute_w_u_vnew_from_h_pallas | 285.0 | 418.9 | 261.83 | 91,9% |
chunk_kda_bwd_dAv_kernel | 148.3 | 209.2 | 130.74 | 88,2% |
_fused_dhu_wy_intra_cumsum_pallas_jit | 900.1 | 951.8 | 594.9 | 66,1% |
Исторически сложилось так, что первые два ядра показали, что простая вычислительная структура и прямой шаблон чтения/записи могут эффективно использовать HBM версии 6e для этой рабочей нагрузки.
В этих измерениях основной целью оптимизации было последнее ядро слияния. Хотя оно исключает большое количество циклов обработки HBM, внутри оно содержит рекуррентность состояний, множественные внутриблочные матричные операции, уменьшение градиента вентиля и обратное суммирование, поэтому оно добавляет вычислительные, конвейерные и локальные накладные расходы.
Измерение: TPU v7x , bf16, varlen seq_lens=[1800,1500,2000,1200,800] (N=5 сегментов), H=16, B=1, T=8192, K=V=128, BT=64; после разбиения на блоки начальное измерение = 2128 = 16×133 (NT=133 = 128 базовых блоков + 5 сегментных дополнений). Пропускная способность HBM принята равной 3,69 ТБ/с.
| Ядро | Время выполнения (мкс) | HBM (MB) | Нижняя граница полосы пропускания (мкс) | использование БВт |
|---|---|---|---|---|
fused_recompute_w_u_vnew_from_h_pallas | 181.675 | 418.9 | 113.52 | 62,5% |
chunk_kda_bwd_dAv_kernel | 100.237 | 209.2 | 56.69 | 56,6% |
_fused_dhu_wy_intra_cumsum_pallas_jit | 784.299 | 951.8 | 257.94 | 32,9% |
4.3 Характеристики производительности контекстного параллельного режима
Измерение: TPU v6e-4 (cp_size=4, 4-ядерный параллелизм SPMD), bf16, один сегмент per_rank_T=8192 → глобальный T=32768, H=16, K=V=128, BT=64; ведущее измерение для каждого ранга после разбиения на блоки = 2128 = 16×133.
Отдельные ядра (с одинаковой базой заполнения):
| Ядро | Время выполнения (мкс) | HBM (MB) | Нижняя граница полосы пропускания (мкс) | использование БВт |
|---|---|---|---|---|
fused_recompute_w_u_vnew_from_h_pallas | 273.4 | 406.3 | 253,95 | 92,9% |
chunk_kda_bwd_dAv_kernel | 146.8 | 202.9 | 126.81 | 86,4% |
_fused_dhu_wy_intra_cumsum_pallas_jit | 872.7 | 915.1 | 571.97 | 65,5% |
CP chunk_gated_delta_rule_bwd_dhu_pre_process | 591.2 | 238.8 | 149.26 | 25,2% |
Измерение: TPU v7x-4 (cp_size=4, 4-ядерный параллелизм SPMD), bf16, один сегмент per_rank_T=8192 → глобальный T=32768, H=16, K=V=128, BT=64; ведущее измерение для каждого ранга после разбиения на блоки = 2128 = 16×133.
| Ядро | Время выполнения (мкс) | HBM (MB) | Нижняя граница полосы пропускания (мкс) | использование БВт |
|---|---|---|---|---|
fused_recompute_w_u_vnew_from_h_pallas | 181 | 406.3 | 110.11 | 60,8% |
chunk_kda_bwd_dAv_kernel | 100 | 202.9 | 54.99 | 55,0% |
_fused_dhu_wy_intra_cumsum_pallas_jit | 785 | 915.1 | 247.99 | 31,6% |
CP chunk_gated_delta_rule_bwd_dhu_pre_process | 511 | 238.8 | 64.72 | 12,7% |
В измеренном пути CP использовались те же три основных ядра, что и в пути без CP; дополнительная работа была выполнена за счет слияния градиентов состояний между рангами.
В историческом контексте основные издержки в рамках CP-процесса приходились не на сам коллектив, а на локальную предварительную обработку информации перед коммуникацией:
- Для предварительной обработки необходимо накопить градиент входного состояния и матрицу обратных переходов вдоль последовательности фрагментов данного ранга;
- Этот этап имеет характер сканирования/уменьшения и его сложно полностью распараллелить, как обычный алгоритм умножения матрицы на фрагмент;
- Сама по себе
all_gatherотносительно недорога, поскольку обмениваются сжатый градиент состояния и матрица переходов, а не тензоры токенов полной последовательности.
Следовательно, оптимизация пути CP должна быть сосредоточена на предварительной обработке:
- уменьшить количество повторных считываний полного фрагмента входных данных;
- улучшить распараллеливание цепочки продуктов/сканирования;
- или переместить/объединить часть логики слияния состояний в существующий этап обратного хода.
4.4 Числовая согласованность
Текущая реализация compute_intra_backward(...) явно преобразует dg_acc + dg_intra к входному эталонному типу данных, а затем обратно к fp32 перед compute_reverse_cumsum_dg(...) . Эта граница типа данных является частью реализованного численного поведения и должна быть сохранена или намеренно перепроверена при рефакторинге слияния. При сравнении с не слиянным эталоном или эталоном XLA следует использовать допуски, соответствующие этой точке усечения.
4.5 Матрица повторной проверки для текущей реализации
Поскольку приведенные выше измерения носят исторический характер, проверка работоспособности головки блока цилиндров должна проводиться отдельно для оценки правильности и производительности. Как минимум, повторная проверка правильности должна включать в себя:
- сохраненный путь
h(rematerialize_for_backward=False) и путь полной рематериализации (rematerialize_for_backward=True); - входные данные фиксированной и переменной длины, включая несколько сегментов и
B>1; - заданное
initial_state,output_final_state=Trueи ненулевой котангенс конечного состояния; - предварительно вычисленные вентили и
use_gate_in_kernel=True, сdelta_time_biasи без него; -
use_qk_l2norm=True; - Несогласованные по контракту
K/Vформы, не относящиеся к CP; - Выполнение CP-операций при поддержке 128-выровненных форм
K/Vи нескольких размеров CP-операций; - Входные данные bf16 и поддерживаемый путь fp32.
При повторной проверке производительности следует указывать точный коммит, поколение и топологию TPU, версии программного обеспечения, формы, распределение переменных, размер CP, политику прогрева, количество итераций, а также указывать, относится ли каждое значение к задержке на устройстве или к сквозной задержке. Оценки пропускной способности на уровне ядра должны сопровождаться сквозным измерением, чтобы были видны накладные расходы на запуск, коллективные и адаптерные затраты.
5. Перспективы оптимизации в будущем
Текущая реализация уже ориентирована на сокращение количества обращений к HBM за счет слияния ядер, но результаты профилирования указывают на два четких направления дальнейших исследований.
5.1 Улучшение производительности версии 7
Согласно историческим измерениям, основные ядра работали быстрее на версии v7x, но при этом демонстрировали более низкое расчетное использование полосы пропускания, чем на версии v6e. Текущий профиль должен сначала подтвердить этот результат; если он останется верным, дальнейшая работа должна быть сосредоточена на том, чтобы унифицированные ядра лучше соответствовали модели выполнения v7:
- В версии 7 следует пересмотреть размеры плиток, выбор мини-пакетов и нагрузку VMEM;
- уменьшить локальные накладные расходы на компоновку и заполнение трафика там, где это хорошо видно в профилях;
- Определите, какие части основного ядра обработки данных ограничены вычислительными ресурсами или производительностью конвейера, а не пропускной способностью.
Цель состоит не в изменении математического разложения, а в перенастройке существующего объединенного пути таким образом, чтобы он лучше масштабировался с более высокой пропускной способностью и вычислительными возможностями версии 7.
5.2 Оптимизация контекстного параллелизма
В пути CP перед обменом данными между рангами добавляется локальная предварительная обработка. Исторический анализ показывает, что эта локальная работа по сканированию/сокращению представляла собой более серьезное узкое место, чем сама коллективная коммуникация; текущий анализ должен подтвердить баланс до начала работы по оптимизации.
Поэтому в будущем при оптимизации CP следует отдавать приоритет следующим аспектам:
- сокращение количества повторных прочтений на этапе предварительной обработки;
- улучшение параллелизации произведения цепочек переходных матриц;
- Изучение возможности интеграции части процесса подготовки градиента состояния CP в существующие обратные ядра.
Схема обмена данными должна оставаться компактной: следует обмениваться сводками на уровне состояний, а не полными тензорами на уровне токенов.