CryptoPatrick / code

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 grad traces 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.