Code · Machine learning
jax-rs
NumPy-style arrays and automatic differentiation in Rust, with a WebGPU backend.
What it is
JAX is Google’s library for numerical computing: NumPy-style arrays plus transformations like grad that compute derivatives automatically. jax-rs brings that model to Rust. It has an Array type with broadcasting and NumPy-like operations, reverse-mode automatic differentiation by tracing computations into a graph, kernel fusion, and a WebGPU backend for running on the GPU, including in the browser.
Using it
use jax_rs::{Array, Shape};
let x = Array::from_vec(vec![1.0, 2.0, 3.0, 4.0], Shape::new(vec![2, 2]));
let y = Array::from_vec(vec![5.0, 6.0, 7.0, 8.0], Shape::new(vec![2, 2]));
let sum = x.add(&y); // element-wise
let prod = x.matmul(&y); // matrix product
let mean = x.mean_all(); // 2.5
Derivatives:
use jax_rs::{Array, Shape, grad};
// f(x) = sum(x² + 2x + 1), so df/dx = 2x + 2
let f = |x: &Array| x.mul(x)
.add(&x.mul(&Array::full(2.0, x.shape().clone(), x.dtype())))
.add(&Array::ones(x.shape().clone(), x.dtype()))
.sum_all_array();
let df = grad(f);
let x = Array::from_vec(vec![1.0, 2.0, 3.0], Shape::new(vec![3]));
println!("{:?}", df(&x).to_vec()); // [4.0, 6.0, 8.0]
Ideas for using it
- Learn autodiff from the inside: read how
gradtraces a function and walks the graph backwards. - Small models in the browser: train or run a tiny network on WebGPU from a Rust/WASM page.
- Optimisation demos: gradient descent on a function you can see, for teaching.
Status
Prototype. The operations and gradients are covered by 547 test functions; performance and coverage against NumPy have not been benchmarked systematically.