Skip to main content

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, and grad.

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

neojax_operators-0.0.6.tar.gz (333.9 kB view details)

Uploaded Source

Built Distribution

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

neojax_operators-0.0.6-py3-none-any.whl (19.7 kB view details)

Uploaded Python 3

File details

Details for the file neojax_operators-0.0.6.tar.gz.

File metadata

  • Download URL: neojax_operators-0.0.6.tar.gz
  • Upload date:
  • Size: 333.9 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

Hashes for neojax_operators-0.0.6.tar.gz
Algorithm Hash digest
SHA256 e671c3cdc24d696e21cdaa4cf5897e5e832fba2b76ca0dc060e46aa36ce1f2dc
MD5 371c19234ca13add91e73710f1a51fb9
BLAKE2b-256 ed33bdb42ada6a3cb5a5eb1ca892be8a4473e1f1857776a1cc0165fb2ae0a288

See more details on using hashes here.

File details

Details for the file neojax_operators-0.0.6-py3-none-any.whl.

File metadata

  • Download URL: neojax_operators-0.0.6-py3-none-any.whl
  • Upload date:
  • Size: 19.7 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

Hashes for neojax_operators-0.0.6-py3-none-any.whl
Algorithm Hash digest
SHA256 b6d2f1c2320c35f05ecf72694ae8bfbbe417184057f5ba4b82f7785ac9ffdc46
MD5 76c78b998424fc816da9c98a9574d097
BLAKE2b-256 e69e8ad91b8a574534d2ab63b6260dfdca297b4dabdadf61088842885c6d9528

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page