Дабы всё можно было реализовать одним кернелом без обращения к глобальной памяти, надо расписать все вычисления так, чтобы нужные активации всегда лежали под рукой у нужного CTA (Cooperative Thread Array). А ещё хочется избавиться от лишних синхронизаций.
Есть две технические сложности:
1. В наивной реализации сначала идёт gate-проекция, а затем up, и CTA нужно тянуться далеко за нужными элементами матрицы весов.
2. Чтобы получить выход down-проекции, необходимо просуммировать по всей промежуточной размерности.
Проблема 1 решается просто — будем перемежать куски up- и gate-проекции, чтобы CTA мог достать их разом из одного места.
Проблема 2 решается гораздо сложнее, и ей посвящена, по сути, большая часть блога. Если коротко — так расписывается то, как данные подаются в тензорные ядра, синхронизируются consumer- и producer-потоки, пайплайнинг и заумный эпилог. Многие штуки зиждутся на Blackwell-специфичных инструкциях, позволяющих перекрывать в одной операции вычисления и трансфер памяти.
⚙️ Есть у них как bf16-, так и mxfp8-реализация.
📊 Замеры
Сравниваются с наивной торчовой реализацией (дико медленной, само собой, и это даже несерьёзно) и
torch.grouped_mm.В bf16 получают ускорение порядка 20% против
grouped_mm и около 2 раз для mxfp8 (против торчового mxfp8). Если убрать requant из торчовой реализации mxfp8, то ускорения почти нет.🤔 Сравнений с другими эффективными fused-кернелами — SonicMoE и Mixture-of-Kittens — нет, так что непонятно, насколько это круто.