Когда стандартный стек машинного обучения трещит под весом кастомной логики, у инженера есть два пути: смириться с ограничениями или случайно написать собственный компилятор (ведь «работает на моей машине» — это уже почти готовый деплой). Если прямо сейчас вы бьетесь с медленными графами XLA и ищете способ выжать максимум из железа, эта история о непредвиденном погружении в мир LLVM укажет вам неожиданный выход.

Введение: От легкого прототипа к тяжелой артиллерии компиляторного строения

Мир современной машинной обработки данных полнó парадоксов. Иногда амбициозные исследовательские задачи решаются не с помощью сложных проектных комитетов, многостраничных RFC-документов и строгого планирования по Agile, а в результате череды экспериментов, динамических языков и банального стремления заставить код работать быстрее здесь и сейчас. Именно такая история приключилась с нашей командой, когда мы исследовали возможности оптимизации вычислений в экосистеме Jax.

Jax прочно закрепился в мире AI и scientific computing как стандарт де-факто для автоматического дифференцирования и функционального программирования на тензорах. Используя связку NumPy-подобного API с XLA (Accelerated Linear Algebra) под капотом, Jax позволяет инженерам забыть о ручной оптимизации CUDA-ядер. Однако в ходе работы над кастомным бэкендом для гибридных вычислительных систем мы столкнулись с ограничениями стандартного графа XLA. Нам потребовалось промежуточное представление, которое давало бы абсолютный контроль над генерацией машинного кода. И вот тут-то всё пошло «не по плану».

Пытаясь написать небольшой вспомогательный транслятор для внутренних нужд (чтобы просто не открывать Stack Overflow в сотый раз за день), мы неожиданно для самих себя спроектировали и реализовали полноценный LLVM-компилятор для Jax. В этой статье мы подробно разберем предпосылки этого архитектурного сдвига, архитектуру получившегося решения, практические примеры интеграции и те неожиданные дивиденды, которые мы получили в производительности.

Но прежде чем погружаться в дебри генерации байт-кода, давайте посмотрим, где именно стандартные абстракции начинают давать трещину под нагрузкой.

Анатомия проблемы: почему стандартный XLA пасует перед кастомной логикой

Чтобы понять, почему возникла необходимость в альтернативном подходе, нужно вспомнить, как устроен классический пайплайн выполнения Jax-программы. Когда вы вызываете функцию, помеченную декоратором @jit, происходит следующий цикл:

  • Трассировка кода (Tracing) с помощью абстрактных значений (ShapedArray), в результате которой строится Jaxpr (промежуточное выражение Jax).
  • Преобразование Jaxpr в StableHLO — высокоуровневый диалект операций машинного обучения.
  • Передача StableHLO в компилятор XLA, который оптимизирует граф, выполняет слияние операций (kernel fusion) и генерирует низкоуровневый код (PTX для GPU или LLVM IR для CPU).

Этот пайплайн невероятно эффективен для плотных матричных умножений и стандартных слоев нейросетей. Но что происходит, когда ваша модель содержит сложную динамическую логику управления потоком, разреженные структуры данных или нестандартные побитовые операции, которые плохо ложатся на матричную парадигму XLA? Здесь разработчики часто упираются в стену.

Ограничения графовых оптимизаций

Мы разрабатывали вычислительный модуль для обработки графов высокой размерности с динамической топологией. XLA сопротивлялся: граф постоянно перестраивался, оптимизационные проходы компиляции занимали больше времени, чем само выполнение инференса (Kubernetes в такие моменты начинает нервно пересчитывать реплики), а отладка сгенерированных ядер превращалась в гадание по логам. Нам требовался инструмент точечного контроля над генерацией машинного кода без избыточной абстракции фреймворков глубокого обучения.

Осознав тупиковость стандартного подхода, мы решились на шаг, который обычно откладывают на самый крайний случай — спуститься на уровень ниже и заговорить с железом на языке компиляторов.

Рождение тран