Better MoE model inference with Warp Decode
🔗 Блогпост
Ребята из Cursor не остановились на Mixture-of-Kittens и реализовали ещё одну примечательную оптимизацию MoE для low-batch-инференса под названием Warp Decode.
Традиционные пайплайны инференса MoE expert-centric: они собирают токены для каждого эксперта, прогоняют вычисления и переставляют их обратно в исходном порядке. Операции перестановок занимают нетривиальное время и существенно замедляют инференс.
Курсоровцы же предлагают параллелизовать не по экспертам, а по выходам. Каждый варп отвечает за одно выходное значение. Инференс реализован через два fused-кернела — gate+up и down. Варп достаёт в потоковом режиме нужную строчку из матрицы весов и проводит операции.
⚡ Так как варпы работают независимо, то всё выходит embarrassingly parallel: вообще не нужно париться по поводу банковских конфликтов, барьеров и синхронизаций. Все редукции выполняются через warp-level-инструкции вида __shfl_xor_sync.
🛠️ Ещё из полезных плюшек стоит отметить следующее:
- 📐 Не нужно паддить до какой-то степени двойки (типа 128).
- 🗂️ Можно избавиться от scatter и combine: токены последовательности раздаются экспертам, а потом всё собирается. Также исчезает необходимость в промежуточных буферах.
🚀 В итоге оно даёт ускорение порядка 1,8× на B200.
💡 А ещё они избавляются от MXFP8-квантизации, ибо и так всё работает достаточно быстро)
📈 На батче из 32 удаётся достичь до 58% максимально достижимой пропускной способности памяти.
Однако авторы утверждают, что их подход не полностью вытесняет expert-centric-исполнение, особенно в сценарии низкой загрузки. На больших батчах оверхеды от перестановок/перегруппировок токенов не так сильно болят.
🔗 Блогпост
Ребята из Cursor не остановились на Mixture-of-Kittens и реализовали ещё одну примечательную оптимизацию MoE для low-batch-инференса под названием Warp Decode.
Традиционные пайплайны инференса MoE expert-centric: они собирают токены для каждого эксперта, прогоняют вычисления и переставляют их обратно в исходном порядке. Операции перестановок занимают нетривиальное время и существенно замедляют инференс.
Курсоровцы же предлагают параллелизовать не по экспертам, а по выходам. Каждый варп отвечает за одно выходное значение. Инференс реализован через два fused-кернела — gate+up и down. Варп достаёт в потоковом режиме нужную строчку из матрицы весов и проводит операции.
⚡ Так как варпы работают независимо, то всё выходит embarrassingly parallel: вообще не нужно париться по поводу банковских конфликтов, барьеров и синхронизаций. Все редукции выполняются через warp-level-инструкции вида __shfl_xor_sync.
🛠️ Ещё из полезных плюшек стоит отметить следующее:
- 📐 Не нужно паддить до какой-то степени двойки (типа 128).
- 🗂️ Можно избавиться от scatter и combine: токены последовательности раздаются экспертам, а потом всё собирается. Также исчезает необходимость в промежуточных буферах.
🚀 В итоге оно даёт ускорение порядка 1,8× на B200.
💡 А ещё они избавляются от MXFP8-квантизации, ибо и так всё работает достаточно быстро)
📈 На батче из 32 удаётся достичь до 58% максимально достижимой пропускной способности памяти.
Однако авторы утверждают, что их подход не полностью вытесняет expert-centric-исполнение, особенно в сценарии низкой загрузки. На больших батчах оверхеды от перестановок/перегруппировок токенов не так сильно болят.