Compile PyTorch models into TensorFlow custom ops backed by Triton PTX kernels
Project description
A compiler that converts PyTorch models into deployable TensorFlow custom ops backed by Triton PTX kernels.
Mirai takes a decorated PyTorch function, runs torch.compile to generate optimized Triton kernels, extracts the PTX, and wraps everything into a TensorFlow custom op — with a single function call.
Quick Start
import torch
import mirai
@mirai.op(name="Pffn")
def pffn(inputs, w_gate, b_gate, w_up, b_up, w_down, b_down):
inputs_t = inputs.transpose(0, 1)
gates = torch.bmm(inputs_t, w_gate) + b_gate.unsqueeze(1)
gates = torch.nn.functional.silu(gates)
vals = torch.bmm(inputs_t, w_up) + b_up.unsqueeze(1)
outputs = torch.bmm(gates * vals, w_down) + b_down.unsqueeze(1)
return outputs.transpose(0, 1)
# One call: torch.compile → PTX extraction → C++ codegen → build script
mirai.build(pffn, sample_inputs=[inputs, w_gate, b_gate, w_up, b_up, w_down, b_down])
Output in ./generated/:
generated/
├── PffnFwd.cc # TF custom op C++ source (forward)
├── PffnBwd.cc # TF custom op C++ source (backward)
├── build.sh # g++ compilation script
├── pffn_api.py # TF Python API wrapper
└── tf32/ # PTX kernels and metadata
├── PffnFwd/
└── PffnBwd/
Build the op:
cd generated && bash build.sh
Dynamic Shapes
Use dynamic=True to generate shape-generic ops that accept variable batch sizes at runtime:
mirai.build(pffn, sample_inputs=sample_inputs, dynamic=True)
A single compiled op then handles any batch size — no recompilation needed:
# TF side: same op binary, different batch sizes
out = pffn_op(inputs_2000, ...) # bs=2000
out = pffn_op(inputs_8000, ...) # bs=8000
See examples/pffn_dynamic.py for a complete example.
Requirements
The pipeline spans two environments (typically on different machines):
| Stage | Dependencies |
|---|---|
Code generation (mirai.build) |
PyTorch 2.x, CUDA |
Op compilation (build.sh) |
TensorFlow 1.x, CUDA |
Run mirai.build() in the PyTorch environment, copy generated/ to the TF environment, then bash build.sh.
Install
pip install mirai-compiler
How It Works
Mirai bridges two frameworks at the PTX level — it leverages PyTorch's compiler stack to produce hardware-optimized GPU kernels, then wraps them as native TensorFlow ops.
PyTorch environment TF environment
┌─────────────────────────────────────────┐ ┌──────────────────────┐
@mirai.op → │ torch.compile → Triton → PTX/CUBIN │ → │ C++ TF custom op │
function │ (max_autotune, fwd + bwd) │ │ (.so, Python API) │
└─────────────────────────────────────────┘ └──────────────────────┘
Under the hood:
-
Trace & Optimize —
torch.compilewithmax_autotuneexplores kernel configurations across the search space. Both forward and backward graphs are traced and optimized independently. -
Intercept & Extract — Mirai rewrites the inductor-generated Python via AST transformers, injecting hooks that capture PTX binaries and tensor metadata at each kernel launch site. The patched code runs in an isolated subprocess.
-
Codegen — Each captured kernel is rendered into a self-contained C++ TF op via Jinja2 templates, complete with shape inference, PTX loading, and CUDA launch logic. A Python API wrapper with
tf.custom_gradientconnects forward and backward ops seamlessly.
Project Structure
mirai/
├── build.py # mirai.build() entry point
├── decorator.py # @mirai.op decorator
├── pipeline.py # Per-kernel: AST patch → PTX extraction → C++ render
├── codegen/ # AST transformers, C++ / build.sh / API renderers
└── templates/ # Jinja2 templates (kernel.cc, build.sh, api.py)
examples/
├── pffn.py # Static shape example
└── pffn_dynamic.py # Dynamic shape example
Development
pip install -e ".[dev]"
black --check mirai/ tests/ examples/
pytest tests/ -v
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file mirai_compiler-0.2.6.tar.gz.
File metadata
- Download URL: mirai_compiler-0.2.6.tar.gz
- Upload date:
- Size: 32.7 kB
- Tags: Source
- Uploaded using Trusted Publishing? Yes
- Uploaded via: twine/6.1.0 CPython/3.13.12
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
dc27f564b36e08b8ab73e3a0fbbe5e9ab6632dff116cfd044f05d0e54500b6b4
|
|
| MD5 |
78b1454a808eb7c14f3709ec176ae287
|
|
| BLAKE2b-256 |
34e8ebf856e9f5d8814a3771e396bf6e3308fbe4604521f7bdaa6ce91f9678a6
|
Provenance
The following attestation bundles were made for mirai_compiler-0.2.6.tar.gz:
Publisher:
publish.yml on dawnop/mirai
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
mirai_compiler-0.2.6.tar.gz -
Subject digest:
dc27f564b36e08b8ab73e3a0fbbe5e9ab6632dff116cfd044f05d0e54500b6b4 - Sigstore transparency entry: 1369859024
- Sigstore integration time:
-
Permalink:
dawnop/mirai@f860451c6223c4989a1412fbfa806d28c11df992 -
Branch / Tag:
refs/tags/v0.2.6 - Owner: https://github.com/dawnop
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish.yml@f860451c6223c4989a1412fbfa806d28c11df992 -
Trigger Event:
push
-
Statement type:
File details
Details for the file mirai_compiler-0.2.6-py3-none-any.whl.
File metadata
- Download URL: mirai_compiler-0.2.6-py3-none-any.whl
- Upload date:
- Size: 32.8 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? Yes
- Uploaded via: twine/6.1.0 CPython/3.13.12
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
7aff2be145d9d8b8bbbb01eb4b4fb044387da5d65f18be13370c74eab7c12bf5
|
|
| MD5 |
12794c266ee113f07273aefc6fa93436
|
|
| BLAKE2b-256 |
289c7ca51353b81afe6bc2887109b2bd22b1a821c62c575b3234eb4228301fc9
|
Provenance
The following attestation bundles were made for mirai_compiler-0.2.6-py3-none-any.whl:
Publisher:
publish.yml on dawnop/mirai
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
mirai_compiler-0.2.6-py3-none-any.whl -
Subject digest:
7aff2be145d9d8b8bbbb01eb4b4fb044387da5d65f18be13370c74eab7c12bf5 - Sigstore transparency entry: 1369859152
- Sigstore integration time:
-
Permalink:
dawnop/mirai@f860451c6223c4989a1412fbfa806d28c11df992 -
Branch / Tag:
refs/tags/v0.2.6 - Owner: https://github.com/dawnop
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish.yml@f860451c6223c4989a1412fbfa806d28c11df992 -
Trigger Event:
push
-
Statement type: