Multi-GPU by default
Data-parallel training via jax.shard_map with exact gradient averaging — the same script runs on 1 or N GPUs, bit-for-bit consistent.
A lean JAX library for physics-informed neural networks.
Define your PDE residual, pick a config, and train:
from jaxpi.models import ForwardIVP, create_model
from jaxpi.samplers import UniformSampler
from jaxpi.training import train
class Burgers(ForwardIVP):
def r_net(self, params, t, x):
u = self.neural_net(params, t, x)
u_t = grad(self.neural_net, argnums=1)(params, t, x)
u_x = grad(self.neural_net, argnums=2)(params, t, x)
u_xx = grad(grad(self.neural_net, argnums=2), argnums=2)(params, t, x)
return u_t + u * u_x - 0.01 / jnp.pi * u_xx
model = create_model(config, Burgers, u0=u0, t_star=t_star, x_star=x_star)
train(config, model, UniformSampler(dom, config.training.batch_size))