Sonata

FlashKDA: ядра Kimi Delta Attention на CUTLASS

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

Загрузка страницы