FlashKDA — набор высокопроизводительных ядер Kimi Delta Attention, построенных на CUTLASS. Проект ориентирован на видеокарты архитектуры SM90 и новее и требует CUDA 12.9+ и PyTorch 2.4+.
По умолчанию сборка определяет текущее CUDA-устройство и компилирует код под его архитектуру. Для сборки колёс или CI-пайплайнов список поддерживаемых архитектур указывается явно, в том числе перечислением через запятую.
После установки FlashKDA подключается автоматически из библиотеки flash-linear-attention — детали интеграции описаны в соответствующем issue проекта. Отключить ядра можно переменной окружения, вернувшись к пути на Triton, а отдельный флаг показывает попадание или промах диспетчеризации.
API принимает тензоры запроса, ключа, значения, гейта до активации и бета-логитов в bf16, скалярный коэффициент масштабирования, log-gate параметры и смещение в fp32, а также нижнюю границу гейта в диапазоне от -5.0 до 0.
Начальное и финальное рекуррентные состояния опциональны и принимают bf16, fp32 либо отсутствие значения; при одновременной передаче их типы должны совпадать. Для батчинга последовательностей переменной длины используется массив кумулятивных длин.
Корректность проверяется тестами на точное совпадение с эталонной реализацией на torch.
Для настройки IntelliSense (clangd) в CUDA/C++ исходниках предусмотрен отдельный скрипт: он генерирует файл с корректными путями репозитория и устанавливает глобальный конфиг clangd.
GitHub ★ 884
