Showing 29 of 137 projects
A JAX-based differentiable spectral modeling library for exoplanets, brown dwarfs, and M dwarfs.
Kernex extends JAX with kmap and kscan for differentiable stencil computations, enabling efficient array transformations.
A JAX + Flax implementation of physics-inspired graph neural networks for solving combinatorial optimization problems like Max-Cut and Maximum Independent Set.
A photovoltaic simulator with automatic differentiation for solar cell modeling and optimization, built on JAX.
A JAX-based library for loopy belief propagation on discrete factor graphs, enabling efficient probabilistic inference.
A differentiable hydrodynamics and magnetohydrodynamics code for astrophysics built with JAX, enabling gradient-based inverse modeling and multi-GPU simulations.
A differentiable ray tracing toolbox for radio propagation simulations, built on JAX for optimization and machine learning.
GPU/TPU accelerated nonlinear least-squares curve fitting using JAX, designed as a drop-in replacement for SciPy's curve_fit.
Serialize JAX, Flax, Haiku, and Objax model parameters using the safetensors format for safer, pickle-free storage.
Composable kernels for scikit-learn implemented in JAX, enabling faster kernel computations and automatic differentiation.
A collection of neural network models ported from torchvision for use with JAX and Flax.
A JAX-based library for fast, composable image augmentation with geometric and color transformations.
A JAX-based library for coreset algorithms that reduce large datasets to smaller representative subsets while preserving statistical properties.
A Python library for accelerated fluid-structure interaction simulations using the immersed boundary lattice Boltzmann method, powered by JAX.
A simple, easy-to-understand library for diffusion models built with Flax and JAX, focusing on readability and learning.
A high-performance JAX-based library for computing optical properties of multilayer thin-film structures using the transfer matrix method.
A 1D3V particle-in-cell simulation code for plasma physics, accelerated using JAX for performance on CPUs, GPUs, and TPUs.
A JAX/Flax implementation of the Fréchet Inception Distance (FID) metric for evaluating generative models.
Flax (JAX) and PyTorch implementations of the DeepSeek-R1-Distill-Qwen-1.5B language model with weight conversion utilities.
A tutorial on building and training a convolutional neural network for MNIST image classification using Flax Linen and Optax.
A Python library for state-based transformation systems in brain modeling and simulation.
MBIRJAX is a Python package for Model Based Iterative Reconstruction (MBIR) of images from tomographic data.
A physical units and unit-aware mathematical system built on JAX for brain dynamics modeling and general scientific computing.
A unified interface for modeling single- and multi-compartment Hodgkin-Huxley neuron models in JAX.
A minimal JAX/Flax implementation of DETR with optimizations like Flash Attention and Sinkhorn solver.
A Python library that uses Taichi Lang to create custom, high-performance operators for brain dynamics simulations.
An Agent-Based Modelling framework built on JAX for scalable and efficient simulations with automatic vectorization and JIT compilation.
A JAX/Flax implementation of the RAFT optical flow estimator with ported checkpoints and reproducible results.
JAX/Haiku implementation of a two-player game approach to learning optimal multi-bidder, multi-item auctions.
Open-Awesome is built by the community, for the community. Submit a project, suggest an awesome list, or help improve the catalog on GitHub.