Catalyst случайно стал компилятором JAX в LLVM для периферийных устройств
Команда разрабатывает Catalyst, компилятор для квантовой программной библиотеки PennyLane. Цель проекта, оптимизировать большие гибридные квантово-классические вычисления через промежуточное представление MLIR. Для захвата классической части (Python, NumPy, SciPy) выбрали JAX: он умеет трассировать Python-функции в вычислительный граф, поддерживает нужные API и уже умеет опускаться до MLIR. Заодно доработали поддержку нативных Python-конструкций управления потоком (while, if, for) и массивов с динамической формой.
Подключая JAX к квантовому конвейеру, разработчики заметили: декоратору @qjit не обязательно передавать квантовые инструкции. Если скормить ему чистый JAX/NumPy-код на обычном Python, Catalyst полностью пропускает бэкенд XLA и опускает представление JAX прямиком в стандартные диалекты MLIR (linalg, arith, scf), а оттуда, в машинный код через LLVM, с обратным распространением через библиотеку Enzyme. Так, работая над квантовым компилятором, команда случайно собрала независимый компилятор из JAX в LLVM.
Стандартный путь JAX через XLA устроен иначе: XLA берёт представление JAXpr, переводит его в диалект StableHLO, прогоняет собственные тяжёлые оптимизации под GPU и TPU и лишь затем использует LLVM для генерации машинного кода под CPU/GPU (для TPU, напрямую). Catalyst перехватывает представление на этапе StableHLO и ведёт его в общие диалекты MLIR, полностью убирая рантайм XLA из цепочки.
Авторы честно признают: пока не уверены, зачем это нужно за пределами их квантовой задачи. Но обход XLA открывает несколько практических возможностей: автономные AOT-бинарники (Ahead-of-Time), не требующие тяжёлого JIT-рантайма PJRT; поддержка массивов с динамической формой без перекомпиляции, к которой обычно принуждает жёсткий XLA; нативные Python-циклы, условия и индексация NumPy «из коробки» через форк Autograph под названием Diastatic Malt; развёртывание на периферийных устройствах, FPGA и нестандартных чипах, куда невозможно принести тяжёлый ML-рантайм; и возможность встраивать собственные диалекты и проходы MLIR до того, как код доберётся до железа.
При этом команда прямо оговаривает: для стандартных задач глубокого обучения и обучения больших трансформеров этот путь XLA не заменит, там по-прежнему нужен обычный JAX. Разработчики обращаются к сообществу с вопросом: полезна ли эта возможность тем, кто работает с нестандартным железом, периферийным ML или инфраструктурой компиляторов.
Ключевые факты
- Команда, строящая квантовый компилятор Catalyst для библиотеки PennyLane, обнаружила: декоратору @qjit можно передать чистый JAX/NumPy-код без единой квантовой инструкции
- В этом случае Catalyst полностью обходит бэкенд XLA и опускает представление JAX напрямую в стандартные диалекты MLIR (linalg, arith, scf), а затем компилирует через LLVM
- Обратное распространение (backprop) обеспечивает библиотека Enzyme, работающая прямо с LLVM IR
- Обход XLA открывает автономные AOT-бинарники без тяжёлого рантайма PJRT, массивы динамической формы без перекомпиляции, нативные Python-циклы и условия, а также развёртывание на периферийных устройствах и нестандартном железе
- Авторы прямо предупреждают: для обучения больших трансформеров и стандартных задач глубокого обучения XLA этот путь не заменит
Почему это важно
Обычно единственный путь из JAX в машинный код лежит через XLA, рантайм, заточенный под GPU и TPU. Команда Catalyst показала, что представление JAX можно опустить напрямую в стандартный MLIR и LLVM, минуя XLA целиком. Это отдельный, независимый от XLA способ компилировать JAX-код, и он получился побочным продуктом совсем другой задачи, построения квантового компилятора.
Кому это важно
В первую очередь, разработчикам, которые работают с нестандартным железом: FPGA, специализированными чипами, периферийными (edge) устройствами, куда нельзя принести тяжёлый ML-рантайм вроде PJRT. Также это интересно инженерам инфраструктуры компиляторов, которые хотят встраивать собственные диалекты и проходы MLIR в конвейер JAX, и разработчикам, которым нужны автономные бинарники без внешних ML-зависимостей.
Как это применить
Пакет ставится командой pip install pennylane-catalyst. Чистый JAX/NumPy-код оборачивается декоратором @qjit (с опцией autograph=True для нативных Python-циклов и условий), квантовые инструкции при этом не нужны вовсе. На выходе получается автономный AOT-бинарник, который поддерживает массивы с динамической формой без перекомпиляции и не требует тяжёлого Python/C++-рантайма для запуска.
Можно ли доверять
Это инженерный отчёт от команды, которая реально строит и поддерживает открытый проект Catalyst для PennyLane, с конкретным разбором архитектуры (MLIR-диалекты, Enzyme, StableHLO) и рабочим примером кода. При этом сами авторы честно пишут, что не уверены в практической ценности находки за пределами их квантовой задачи: бенчмарков, сравнений производительности с XLA или данных о зрелости этого пути в тексте нет.
Риски и подводные камни
В материале нет ни одной цифры производительности и ни одного сравнения с XLA, насколько зрелый и готовый к продакшену этот обходной путь, неизвестно. Авторы прямо предупреждают: для стандартных задач глубокого обучения и обучения массивных трансформеров этот подход XLA не заменяет, там по-прежнему нужен обычный JAX. Ценность идеи для не-квантовых применений (периферийный ML, кастомное железо) пока остаётся открытым вопросом, а не подтверждённым результатом.
«Мы случайно построили конвейер компиляции MLIR для JAX.»
— команда Catalyst