There are currently 25 open-source projects built with Flax, with a combined total of 30.8k GitHub stars. The most common language among these projects is Python.
Showing 25 open-source projects
Official JAX/Flax implementation of Vision Transformer (ViT) and MLP-Mixer for image recognition, with pre-trained models.
A JAX implementation of OpenAI's Whisper model offering up to 70x faster transcription on TPUs.
A JAX library for rapid prototyping of large-scale attention-based vision models across images, video, audio, and multimodal data.
A JAX/Flax-based framework for easy and scalable pre-training, fine-tuning, evaluation, and serving of large language models.
A high-performance, scalable LLM library and reference implementation written in pure Python/JAX for training on TPUs and GPUs.
Official repository for Big Transfer (BiT) models, providing pre-trained visual representations for efficient transfer learning across computer vision tasks.
JAX (Flax) implementations of reinforcement learning algorithms for continuous action spaces, designed for research.
A low-level Gaussian process framework in JAX and Flax, designed for maximum flexibility and close alignment with mathematical notation.
A JAX library implementing Lie groups for rigid body transformations in computer vision and robotics.
A JAX library for automatically generating equivariant neural network layers for arbitrary symmetry groups via constraint solving.
A collection of pretrained deep learning models (StyleGAN2, GPT2, VGG, ResNet) for the Jax/Flax ecosystem.
A FlashAttention 2 implementation for JAX with block-wise document mask optimization and context parallelism for efficient long-sequence training.
Unofficial JAX/Flax implementations of deep learning research papers for vision transformers and other architectures.
A JAX transform that implements LoRA (Low-Rank Adaptation) for efficient fine-tuning of large models with minimal memory overhead.
Flax implementations and pretrained checkpoints for ResNet, Wide ResNet, ResNeXt, ResNet-D, and ResNeSt in JAX.
Official JAX implementation of XMC-GAN for text-to-image generation using cross-modal contrastive learning.
A JAX + Flax implementation of physics-inspired graph neural networks for solving combinatorial optimization problems like Max-Cut and Maximum Independent Set.
A collection of neural network models ported from torchvision for use with JAX and Flax.
A simple, easy-to-understand library for diffusion models built with Flax and JAX, focusing on readability and learning.
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 minimal JAX/Flax implementation of DETR with optimizations like Flash Attention and Sinkhorn solver.
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.
Open-Awesome is built by the community, for the community. Submit a project, suggest an awesome list, or help improve the catalog on GitHub.