← На главную

Catalyst случайно сделал компилятор JAX на MLIR и Enzyme вместо XLA

30.07.2026 23:46 · hackernews

Команда разработчиков квантового компилятора Catalyst для библиотеки PennyLane случайно собрала ещё и MLIR-компилятор для JAX. Они строили оптимизатор гибридных квантово-классических рабочих процессов на базе MLIR. Для захвата Python-кода выбрали JAX — он умеет трассировать функции, строить вычислительный граф, поддерживает NumPy и SciPy, а ещё умеет опускаться до MLIR. Заодно добавили поддержку нативного Python control flow и динамических форм массивов.

В процессе встройки JAX в квантовый пайплайн обнаружилось: в Catalyst @qjit можно передавать чистый классический JAX-код вообще без квантовых инструкций. Тогда Catalyst пропускает XLA, транслирует представление JAX в стандартный MLIR и компилирует в машинный код через LLVM. Для обратного распространения ошибки используют Enzyme — он работает прямо над LLVM IR. Квантовые градиенты считают отдельные самописные MLIR-проходы.

Так вышло, потому что изначально команда хотела не изобретать классическую компиляторную инфраструктуру, а опереться на зрелый LLVM и добавить квантовую поддержку. Важно сохранять структуру квантовой программы: циклы, if, смешение массивов и квантовых инструкций. Плюс не хотелось ломать квантовое автодифференцирование. Вывод простой: квантовые алгоритмы никогда не бывают только квантовыми. Вокруг куча классической обработки — подготовка входов, разбор выходов, контроль оборудования QPU для коррекции ошибок. Поэтому компилировать нужно весь гибридный граф в один исполняемый файл, который работает максимально близко к железу.

Catalyst перехватывает представление JAX на пути к XLA и вместо традиционного пайплайна опускает его в стандартные MLIR-диалекты — linalg, arith, scf. Дальше — LLVM и Enzyme, XLA runtime выбрасывается. В статье признаются, что в идеале хотели бы вообще отказаться от HLO и захватывать семантику NumPy напрямую, но пока JAX делает за них грязную работу по снижению линейной алгебры.

Это не конкурент XLA для стандартных задач глубокого обучения — там XLA годами оптимизировали под GPU и TPU. Но связка JAX → MLIR → LLVM без XLA даёт неожиданные бонусы. Можно собирать автономные AOT-бинарники, которым не нужен внешний ML-рантайм вроде PJRT. Поддерживаются динамические формы без мучительной рекомпиляции. Нативно работают Python-циклы, if, булевы операции, даже присваивания в NumPy — всё дифференцируемо. Упрощается запуск моделей на edge-устройствах, FPGA или экзотических чипах: чистый LLVM IR без тяжёлого рантайма легко утащить куда угодно. А ещё код попадает в стандартный MLIR, так что в него легко вставлять собственные диалекты и проходы компилятора.

Команда называет это забавным побочным эффектом и не уверена, насколько он полезен за пределами квантовых вычислений. Но если вы работаете с кастомным железом, edge ML или компиляторами — возможно, такой мост между JAX и MLIR вам пригодится.

Читать оригинал →