Skip to main content

TorchBurn Logo

TorchBurn

A Hardware-Agnostic, High-Performance PyTorch Compilation Backend in Rust
Zero-Copy DLPack FFI, BLAKE3 Graph Caching, Single-Pass Kernel Loop Fusion & Multi-Engine Execution.

CI PyPI Python License Supported Engines Zero CUDA


⚡ What is TorchBurn?

TorchBurn bridges PyTorch's compiler frontend (torch.compile) with a high-performance, multi-engine Rust backend. By default, it compiles PyTorch computation graphs into zero-copy, cache-optimized Native CPU kernels using Rayon chunked parallelism and AVX2/NEON SIMD vectorization. It also provides plug-and-play support for Burn's CPU engine (burn_ndarray) and Universal GPU engine (burn_wgpu across Vulkan, DirectX 12, Metal, and WebGPU) with zero CUDA installation required.

Tensors cross the Python ↔ Rust boundary zero-copy using the open DLPack standard, graph DAGs are cached with BLAKE3 structural hashing, and unsupported operators safely fall back to eager PyTorch.

import torch
import torchburn  # Automatically registers the "torchburn" backend

model = torch.nn.Sequential(
    torch.nn.Linear(512, 1024),
    torch.nn.ReLU(),
    torch.nn.Linear(1024, 256),
).eval()

# One line to compile: runs on native CPU by default, or opt into WGPU
compiled_model = torch.compile(model, backend="torchburn")
output = compiled_model(torch.randn(32, 512))

🧠 Universal LLM Engine (Zero CUDA, Zero llama.cpp)

TorchBurn v0.5.5 introduces torchburn.LLM: a high-level, universal language model inference engine that runs any model directly from Hugging Face Hub or local checkpoints with 5–9 lines of code.

  • No CUDA, No llama.cpp, No GGUF conversion: Executes directly on raw PyTorch weights (.safetensors).
  • Hardware Auto-Dispatch: Seamlessly dispatches across all hardware:
    • CPUs: AVX-512 VNNI / AVX2 / ARM NEON with pure-Rust decoders (up to 76.5 tokens/sec).
    • iGPUs & dGPUs: Intel Iris Xe, AMD Radeon, Apple Silicon, and NVIDIA via the End-to-End WGPU Compute Graph Decoder (Vulkan) with 1-shot command stream submission (up to 23.3 tokens/sec on Intel Iris Xe).
  • Hugging Face Hub Integration: Direct loading with automatic token discovery (token="...", HF_TOKEN, or huggingface-cli login cache).
  • Universal Quantization: INT4 SIMD (W4A32), INT8, and FP32.

5-Line Generation

import torchburn as tb

# Load any Hugging Face or local model in 1 line
llm = tb.LLM.from_pretrained("Qwen/Qwen2.5-0.5B-Instruct", quant="int4", device="auto")

# Generate completion in 1 line
print(llm.generate("Explain quantum computing in two sentences."))

Real-Time Streaming

import torchburn as tb

llm = tb.LLM.from_pretrained("Qwen/Qwen2.5-0.5B-Instruct", quant="int4")
for token in llm.stream("Once upon a time in a digital kingdom:"):
    print(token, end="", flush=True)

Interactive Multi-Turn Chat

import torchburn as tb

llm = tb.LLM.from_pretrained("Qwen/Qwen2.5-0.5B-Instruct", quant="int4")
llm.chat(system_prompt="You are a helpful and concise AI.")

Command-Line CLI

# Interactive chat in your terminal
python -m torchburn.llm chat --model Qwen/Qwen2.5-0.5B-Instruct --quant int4

# Single prompt generation
python -m torchburn.llm generate "Explain black holes" --model models/qwen_0_5b --stream

# Hardware benchmark (measures tok/s and latency)
python -m torchburn.llm benchmark --model models/qwen_0_5b --device cpu --tokens 64

🚀 Key Highlights (v0.6.1)

  • 🔥 Peak-Optimization Release (v0.6.1): fat LTO + codegen-units=1 release profile, portable per-target baselines (x86-64-v2 / neoverse-n1 / apple-m1 / apple-a14) with runtime AVX2/AVX-512-VNNI/NEON dispatch, cached CPUID (no per-row CPUID), vectorized m==1 GEMV (f32x8 FMA), lock-optimized BLAKE3 cache, per-bucket memory pooling, and allocation-free contiguity checks. Prebuilt wheels ship for Windows AMD64, Linux x86_64+aarch64, macOS arm64+x86_64 (plus CUDA and native variants) — see CHANGELOG.
  • Native CPU by Default: Out-of-the-box zero-copy execution on CPU with zero GPU setup or shader compilation delays. Reaches 98.2% parity with Intel MKL on $1024^3$ GEMM (12.29 ms vs 12.08 ms), with -48.5% GEMM improvement in v0.5.5.
  • 🔄 Single-Pass Kernel Loop Fusion: Fuses multi-node unary/binary DAGs into single memory sweeps with stack-allocated [T; 32] scratch space, eliminating heap allocations in worker threads.
  • 🏎️ Chunked SIMD Parallelization: Rayon L1/L2-aware chunking (PAR_CHUNK = 16 * 1024) with wide f32x8 vectorized polynomials for GELU (7.7× speedup: 9.83 ms → 1.28 ms), sigmoid, tanh, silu, and softmax.
  • 📦 Prepared Graph Pre-Planning: prepare_graph() pre-plans memory slot assignments and fusion plans once, skipping graph traversal and HashMap lookups on every forward pass.
  • 🧮 Parallel Epilogue Fusion: Fuses Linear and GEMM activation epilogues (ReLU, GELU, Sigmoid, SiLU) directly into multi-threaded chunked matrix output writes.
  • 🎮 Multi-Engine Flexibility: Seamlessly toggle between native_cpu, burn_ndarray, and burn_wgpu (Vulkan / DX12 / Metal).
  • 🧬 BLAKE3 Structural Graph Caching: Nanosecond-level cache lookups with LRU promotion bypass re-tracing overhead on warm runs.
  • 🛡️ Safe Eager Fallback: Unrecognized nodes fall back with bounded warnings and op_coverage() telemetry.
  • 🔒 100% Test Passing: Comprehensive test coverage across 450 native operators verified against PyTorch ground truth.
  • 🚀 O(1) Dispatch: HashMap-based operator dispatch replaces 22-module linear scan, eliminating per-node string comparison overhead.
  • 🧵 Rayon-Parallel Kernels: Losses, embedding, matmul backward, softmax, and gradient accumulation all parallelized via rayon for large tensors.
  • 🎯 Zero-Copy View Cloning: Arc<[i64]> slot views eliminate Vec clone overhead on every node execution.

🏛️ Multi-Engine Architecture: Why 3 Engines?

TorchBurn features 3 distinct execution engines tailored for different deployment environments:

  1. native_cpu (Default):
    • Characteristics: Hand-tuned Rust kernels using rayon, matrixmultiply, and wide f32x8 SIMD.
    • Best For: Maximum single-node and server CPU inference throughput, zero dependencies, instant execution.
  2. burn_ndarray (Pure Rust CPU Engine):
    • Why is it needed?:
      • Golden Reference: 100% safe, pure-Rust fallback with zero C/CBLAS dependencies, critical for cross-compilation (e.g., embedded, WASM, musl).
      • Headless CI Stability: Provides a reliable CPU fallback in headless CI runners where virtualized GPU/Metal drivers return uninitialized buffers.
      • Burn Ecosystem Interoperability: Allows direct bridge and graph execution within Burn's native training/deployment pipeline.
  3. burn_wgpu (Universal GPU Acceleration):
    • Characteristics: WebGPU compute shaders executing across AMD, Intel, Apple Silicon, and NVIDIA via Vulkan, DirectX 12, or Metal.
    • Best For: GPU acceleration on consumer hardware without installing multi-gigabyte CUDA toolkits.

📊 Performance Benchmarks (v0.5.5 Native CPU vs Intel MKL / PyTorch Eager)

System: Intel Core i7-11800H @ 2.30 GHz (8 cores / 16 threads), Windows 11 x86_64, FP32

Workload PyTorch Eager (MKL/AVX2) TorchBurn native_cpu Status / Ratio v0.5.5 Improvement
GEMM $1024 \times 1024 \times 1024$ 12.08 ms 13.5 ms 89.5% Parity -48.5% vs v0.5.4 (26.2 ms)
GEMM $256 \times 256 \times 256$ 328 µs -24.7% vs v0.5.4
GELU Activation ($1024^2$) 0.94 ms 1.28 ms 1.36× of eager 7.7× faster vs pre-chunked
Multi-Head Attention ($B=4, H=8, T=128, D=64$) 0.93 ms 1.13 ms 82.2% Parity Zero-copy QKV projection
Linear + Epilogue ($128 \times 512 \to 1024$) 0.35 ms 0.55 ms 1.57× of eager Vectorized parallel epilogue
Softmax ($2048 \times 2048$) 8.12 ms 8.84 ms 91.8% Parity SIMD-accelerated (new in v0.5.5)
Decoder Step (Qwen 0.5B int4) 3.46 ms/token -9.7% vs v0.5.4

⚙️ Device & Engine Selection

TorchBurn executes on native_cpu by default. You can easily inspect or customize execution target via environment variables or Python API:

import torchburn

# Check active execution engine
print(torchburn.active_engine())  # 'native_cpu' (default), 'burn_ndarray', or 'burn_wgpu'

Note: engine selection is process-wide and read at import time — set TORCHBURN_ENGINE before your script starts (see table below).

Environment Variable Controls

Environment Variable Allowed Values Description
TORCHBURN_ENGINE native_cpu (default), burn, burn-wgpu Explicitly select execution backend
TORCHBURN_DEVICE cpu (default), gpu High-level device target switch
TORCHBURN_WGPU_BACKEND vulkan, dx12, metal, gl Force specific graphics API for WGPU

Run with WGPU acceleration:

TORCHBURN_ENGINE=burn-wgpu python your_model.py

🏗️ Architecture & Data Flow

   torch.compile(model, backend="torchburn")
                     │
                     ▼
          torch._dynamo / FX Graph
                     │
                     ▼
       TorchBurn FX Partitioning Engine
        ┌────────────┴────────────┐
        │                         │
        ▼                         ▼
   Supported Nodes         Unsupported Nodes
   (130+ ops)                     │
        │                         ▼
        │                 Safe Eager Fallback
        ▼                 (Ground-truth PyTorch)
   Zero-Copy DLPack FFI (`engine.rs:753` `allow_threads`)
         │
         ▼
     Rust Execution Core (release):
     ├── O(1) HashMap Dispatch `dispatch_op.rs` `OnceLock<HashMap>`
     ├── BLAKE3 LRU Cache `cache.rs:23` 1024 + `pool.rs:34` best-fit MaybeUninit
     ├── L1 16KB + `wide f32x8` SIMD `ops.rs:22` + online softmax `activations.rs:270`
     ├── Rayon-parallel losses/embedding/matmul/softmax/gradient-accum
     ├── OpenBLAS Skylake `blas.rs:7` + `matrixmultiply` tiled GEMM `linalg.rs:1`
     ├── Fusion `fusion.rs:55` `ConvBnRelu` + QKV+softmax+V (v2)
     └── Burn WGPU 16×16 tiled `wgpu_kernels/matmul.wgsl` vec4 shaders LRU `wgpu_backend.rs:179`
        │
        ▼
   Zero-Copy DLPack Output Capsules ──► torch.Tensor

🧩 Supported Operators (450 wired – v0.4.1)

Click to expand full operator matrix (450)
Category Operators Count
Elementwise add, sub, mul, div, neg, reciprocal, abs, sign, clamp, fmod, remainder, bitwise_and/or/xor/not, copysign, ldexp, nextafter, heaviside, isclose, allclose, equal, isreal, is_complex 32
Math & Transcendentals exp, exp2, expm1, log, log2, log10, log1p, sqrt, rsqrt, square, pow, sin, cos, tan, asin, acos, atan, sinh, cosh, tanh, erf, erfc, asinh, acosh, atanh, sinc, i0/i1/i0e/i1e, bessel_j0/j1/y0/y1, digamma, lgamma, polygamma, mvlgamma, erfinv, erfcinv, ndtri, ndtr, log_ndtr, logit, expit, rad2deg, deg2rad, trunc, frac, logspace, eye, diag, triu/tril 58
Activations relu, sigmoid, tanh, gelu, silu, leaky_relu, elu, selu, softplus, mish, softmax, log_softmax, hardtanh, hardsigmoid, glu, celu, hardshrink, softshrink, tanhshrink, threshold, logsigmoid, rrelu, bernoulli, multinomial 24
Linear Algebra linear, matmul, bmm, addmm, dot, t, transpose, mv, vdot, baddbmm, addbmm, addmv, kron, inner, outer, linalg_multi_dot, linalg_vander, linalg_vecdot, linalg_cross, linalg_tensordot, linalg_norm, frobenius_norm, nuclear_norm, matrix_rank, cholesky, qr, svd, eig, lu 32
Reductions sum, mean, max, min, argmax, argmin, std, var, var_mean, std_mean, prod, cumsum, all, any, amax, amin, count_nonzero, nansum, nanmean, nanprod, nanmin, nanmax, nanmedian, cummax, cummin, logcumsumexp, logsumexp, cov, corrcoef 30
Normalization layer_norm, batch_norm, group_norm, rms_norm, instance_norm, local_response_norm, channel_shuffle 7
Shape & Indexing reshape, view, view_as, permute, squeeze, unsqueeze, expand, expand_as, broadcast_to, broadcast_tensors, flatten, cat, stack, split, chunk, vsplit, hsplit, dsplit, tensor_split, unbind, select, narrow, gather, index_select, take_along_dim, index_reduce, scatter_max/min, tile, roll, pixel_shuffle, unfold, fold, pixel_unshuffle, grid_sample, affine_grid, as_strided, empty_strided, take, put, index_fill, masked_select/scatter, index_add/put 45
Convolution & Pooling conv1d, conv2d, conv3d, conv_transpose1d, conv_transpose2d, conv_transpose3d, max_pool1d, max_pool2d, max_pool3d, avg_pool1d, avg_pool2d, avg_pool3d, adaptive_avg/max_pool1d/2d/3d, fractional_max_pool2d/3d, lp_pool1d/2d/3d, max_unpool1d/2d/3d 28
Transformer/LLM scaled_dot_product_attention, flash_attention, fused_swiglu/geglu/rmsnorm_residual, embedding, embedding_bag, rope, multi_head_attention_forward, lstm/gru/rnn_cells 12
Losses mse_loss, huber_loss, smooth_l1_loss, cross_entropy, nll_loss, binary_cross_entropy, kl_div, poisson_nll, margin_ranking, hinge_embedding, soft_margin, cosine_embedding, triplet_margin, ctc_loss, bincount, unique, kthvalue, median, histogram, bucketize, searchsorted, meshgrid 22
Creation/Quant/FFT full, zeros, ones, arange, linspace, rand/randn/randint/randperm, empty, zeros_like, ones_like, full_like, randn_like, rand_like, randint_like, eye, diag, hann/bartlett/blackman/hamming/kaiser/gaussian windows, stft, istft, quantize/dequantize_per_tensor/channel, int8_gemm, nf4_dequantize, fft, ifft, rfft, irfft, fft2, ifft2, fftn, ifftn, fftshift, ifftshift, complex, real, imag, angle, polar, conj 42

See docs/ops_coverage.md for full signatures and test coverage metrics.


🚀 Performance Roadmap

The full phased plan for kernel/backend optimization (CPU int4 GEMV, iGPU latency, CUDA backend, GGUF import, speculative decoding) with measured baselines and benchmark gates lives in docs/OPTIMIZATION_ROADMAP.md.


📦 Wheels (v0.6.1)

Prebuilt portable wheels are published to PyPI on every v* tag (see CHANGELOG.md). Baselines are portable; faster ISA paths (AVX2/AVX-512-VNNI/NEON) are selected at runtime:

Platform Architectures Notes
Windows AMD64 (x86-64-v2) Vulkan/DX12 via burn-wgpu
Linux x86_64 (x86-64-v2), aarch64 (neoverse-n1) Vulkan; CUDA wheel from cuda-build job
macOS 11+ arm64 (apple-m1), x86_64 (x86-64-v2) Metal / Vulkan
Native (self-hosted) host (target-cpu=native) Peak bench wheel, not for PyPI
pip install torchburn  # portable optimized wheel

🧪 Testing & Validation

TorchBurn 450 native ops, 553 tests (test_all_450_ops.py 450 distinct), validate_450.py 48/48 batch4 pass torch.allclose(atol=1e-4):

# Run full 450-op sweep (release)
python -m pytest tests/test_all_450_ops.py -q  # 450 distinct
python -m pytest tests/ -q  # 553 passed, 5 deselected (BertTiny/BenchmarkSuite)

# Validate batch4 48 vs PyTorch
python validate_450.py  # 48/48 PASS

# Force CPU or GPU
TORCHBURN_DEVICE=cpu python -m pytest tests/ -q
TORCHBURN_DEVICE=gpu python bench_full.py  # Iris Xe Vulkan 134ms softmax

# Lints
cargo clippy -- -D warnings
cargo fmt --check

📄 License

TorchBurn is open-source software licensed under the Apache 2.0 License. See LICENSE for details.

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

torchburn-0.6.3.tar.gz (742.6 kB view details)

Uploaded Source

Built Distributions

If you're not sure about the file name format, learn more about wheel file names.

torchburn-0.6.3-cp39-abi3-win_amd64.whl (8.5 MB view details)

Uploaded CPython 3.9+Windows x86-64

torchburn-0.6.3-cp39-abi3-manylinux_2_28_x86_64.whl (5.5 MB view details)

Uploaded CPython 3.9+manylinux: glibc 2.28+ x86-64

torchburn-0.6.3-cp39-abi3-macosx_11_0_x86_64.whl (5.2 MB view details)

Uploaded CPython 3.9+macOS 11.0+ x86-64

torchburn-0.6.3-cp39-abi3-macosx_11_0_arm64.whl (4.9 MB view details)

Uploaded CPython 3.9+macOS 11.0+ ARM64

File details

Details for the file torchburn-0.6.3.tar.gz.

File metadata

  • Download URL: torchburn-0.6.3.tar.gz
  • Upload date:
  • Size: 742.6 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for torchburn-0.6.3.tar.gz
Algorithm Hash digest
SHA256 c04a7abec4d96246cfb7b1388c25bde7b58347bd23cde3e296f87a6a846dc3f9
MD5 f98e5f67dded89fc858c5bb97622bb13
BLAKE2b-256 ab4077a2cb1745950899a9d76193a056890a643e4666c9fa1b6bc4cffaabe16c

See more details on using hashes here.

File details

Details for the file torchburn-0.6.3-cp39-abi3-win_amd64.whl.

File metadata

  • Download URL: torchburn-0.6.3-cp39-abi3-win_amd64.whl
  • Upload date:
  • Size: 8.5 MB
  • Tags: CPython 3.9+, Windows x86-64
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for torchburn-0.6.3-cp39-abi3-win_amd64.whl
Algorithm Hash digest
SHA256 745d0df4776494e58e72fd27d1630e82a1147a0ae3ede4744fb1fe1b2b84f037
MD5 2036cb00320ef3ea1c7ec6e02683328b
BLAKE2b-256 051b37a08dd5b5de60dbd0f5e00f2927e2baa062a5cfa01ddb9a60cb22f27229

See more details on using hashes here.

File details

Details for the file torchburn-0.6.3-cp39-abi3-manylinux_2_28_x86_64.whl.

File metadata

File hashes

Hashes for torchburn-0.6.3-cp39-abi3-manylinux_2_28_x86_64.whl
Algorithm Hash digest
SHA256 0daa4935e54eb3d26d9b9d8cef3e2444c77b9eff0f8f469aed05aada9c4176d3
MD5 64a02dda9097a2a2091444e97082cfe6
BLAKE2b-256 bb93abb8120f87d97e0a0ddb62ddb844e1a46b3029f75bd65b3ea902c12230f1

See more details on using hashes here.

File details

Details for the file torchburn-0.6.3-cp39-abi3-macosx_11_0_x86_64.whl.

File metadata

File hashes

Hashes for torchburn-0.6.3-cp39-abi3-macosx_11_0_x86_64.whl
Algorithm Hash digest
SHA256 69ff8a7d38ef80abb2bbcb87907528144c4e2b6c8f55506d2ce124c34aaf2412
MD5 1da85d0eb0e7c03eb05d63460d8de9a9
BLAKE2b-256 5444af35d443a4a53f5d2d94ad31bd9eb176487c2929b1132e5c559cc1a2b3bd

See more details on using hashes here.

File details

Details for the file torchburn-0.6.3-cp39-abi3-macosx_11_0_arm64.whl.

File metadata

File hashes

Hashes for torchburn-0.6.3-cp39-abi3-macosx_11_0_arm64.whl
Algorithm Hash digest
SHA256 4936784637b247500097f42954b808018031663f5b9630a3af26f6bf668a9e37
MD5 b20f2a657b707bda885efc2418d9ecb6
BLAKE2b-256 4ca4b51d149d225eb12672036702a71cba2e98e66ec58df3688b2bd1968cd0ea

See more details on using hashes here.

Release history Release notifications | RSS feed

This release

0.6.3 This release

5 files

0.5.4

5 files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page