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 Distributions

No source distribution files available for this release.See tutorial on generating distribution archives.

Built Distribution

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

torchburn_cuda-0.6.3-cp39-abi3-manylinux_2_35_x86_64.whl (5.7 MB view details)

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

File details

Details for the file torchburn_cuda-0.6.3-cp39-abi3-manylinux_2_35_x86_64.whl.

File metadata

File hashes

Hashes for torchburn_cuda-0.6.3-cp39-abi3-manylinux_2_35_x86_64.whl
Algorithm Hash digest
SHA256 24e0212ad9fac7cb2f9b64237a00dfe9f42b6c4513dda7b444f114b424a706a3
MD5 90018dc56e50de61f7d26fe6033c0338
BLAKE2b-256 4d4fc6ac5c82cfddd268909134d1f9ff8905f193a5f50cc8d1d7e7eb97aff06c

See more details on using hashes here.

Release history Release notifications | RSS feed

This release

0.6.3 This release

1 file

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