Neural Operators in JAX
Project description
neojax
neojax (Neural Operators in JAX) is an implementation of Neural Operators built on top of JAX and Equinox. It provides a clean, modular API inspired by the original neuraloperator library.
Currently, neojax is in its early stages. Only the Fourier Neural Operator (FNO) is available, with support for:
- 1D, 2D, 3D, 4D, and 5D coordinate-based operator learning.
- Grid-based positional embeddings (
GridEmbeddingNd). - Symmetrical domain padding (
DomainPadding). - Pointwise MLP (Channel-MLP) expansions for improved expressivity.
- Various skip connections (Linear, Soft-Gating, Identity).
- Seamless integration with JAX's
vmap,jit, andgrad.
It is designed to be fully compatible with all JAX features such as vmap, jit, and grad.
Installation
Install the python package via pypi
pip3 install neojax-operators
Quickstart
neojax exposes a similar API to neuraloperators and equinox and should therefore be familiar to use:
import jax.numpy as jnp
from neojax.models import FNO
fno = FNO(n_modes=(64, 64),
hidden_channels=64,
in_channels=2,
out_channels=1
)
x = jnp.ones((2, 64, 64))
pred = fno(x)
For a more in detail introduction refer to the examples in the documentation.
Benchmarks
Coming soon (Please refer to BENCHMARKS.md for a performance comparison of neojax and neuraloperator.)
Motivation
Jax has become ubiquitous in Scientific Machine Learning (SciML) and Scientific Computing. This is largely due to its core design, which embraces mathematical and functional transformations (like jit, vmap, and grad) and seamlessly integrates with NumPy-like paradigms. However, despite Neural Operators fundamentally shaping the SciML landscape and being frequently used for solving PDEs, a comprehensive, native Jax implementation has been notably missing. neojax was created to bridge this gap, bringing the performance, predictability, and ecosystem of Jax to the Neural Operator community.
Design Choices
Although originally conceived as a direct port of the PyTorch neuraloperator library, neojax evolved into a ground-up, Jax-native re-implementation. This approach avoids the pitfalls of forcing PyTorch idioms into a functional framework and significantly reduces internal complexity. By building directly on equinox, neojax aligns perfectly with Jax's pure-functional design principles while maintaining a clean, accessible, and class-based API.
Roadplan
The first and currently only supported Neural Operator is a simple Fourier Neural Operator (FNO). In upcoming releases more models and components will be added.
Contributions
If you'd like to contribute any features, models, or fix implementation errors, please do so. Any contributions are appreciated. Have a look at the CONTRIBUTING.md guide for details on how to do so.
Citation
If you use neojax in your research, please cite it using the following BibTeX entry:
@software{neojax,
author = {Paul Gekeler},
title = {neojax: Neural Operators in Jax},
year = {2026},
url = {https://github.com/paulgekeler/neojax}
}
Project details
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file neojax_operators-0.0.5.tar.gz.
File metadata
- Download URL: neojax_operators-0.0.5.tar.gz
- Upload date:
- Size: 530.7 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: uv/0.11.14 {"installer":{"name":"uv","version":"0.11.14","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"macOS","version":null,"id":null,"libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
d0e422c3eb4e0e14340df649fd5fb9f6246fe6239ff2eaa502dd581753aa0fb0
|
|
| MD5 |
184a0f66b7a88aaca870831770f6049f
|
|
| BLAKE2b-256 |
2fd284082422ad783e04756fe67185f392903670ac057820e7e4e73bc895295f
|
File details
Details for the file neojax_operators-0.0.5-py3-none-any.whl.
File metadata
- Download URL: neojax_operators-0.0.5-py3-none-any.whl
- Upload date:
- Size: 17.6 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: uv/0.11.14 {"installer":{"name":"uv","version":"0.11.14","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"macOS","version":null,"id":null,"libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
d0ac4e10344f46b62d7d1a8d82e769d3b4785fa20da8675bd3e33c6487dda599
|
|
| MD5 |
90c1df992e2689131bf91b98040729af
|
|
| BLAKE2b-256 |
a566156fdc67036b042fde9b1e77eea967f8d3fa1aa3cfa2e909876e9624bc7c
|