jax-ml/jax
68.0
Adequate · 26 September 2026
556.2k
lines of production code
Python
with C++
3
measurements over time
What this system is
JAX is a high-performance numerical computing library that enables automatic differentiation, vectorization, and just-in-time compilation for Python and C++. It provides a NumPy-compatible API alongside specialized modules for linear algebra, sparse operations, and statistical distributions, while supporting execution across CPUs, GPUs, and TPUs. The system includes advanced features for distributed training, custom kernel programming via Pallas, and interoperability with TensorFlow through JAX2TF.
How it got here
2018–2020 — API stabilization and build modernization
39 changes.
This period focused on restructuring the JAX codebase by migrating implementation details to private namespaces and establishing stable public APIs for core modules like numpy, scipy, and lax. The work also involved modernizing the build infrastructure with a new Bazel-based system and introducing the jax2tf package to enable bidirectional interoperability with TensorFlow.
2021–2023 — Pallas, Mosaic, and MLIR expansion
47 changes.
This period focused on integrating the Pallas programming model and Mosaic compiler infrastructure as core components, establishing stable public APIs for MLIR and custom extensions. It also expanded JAX's scientific computing capabilities through new SciPy implementations, experimental sparse matrix support, and a new interactive debugger, while restructuring the build system and backend plugins for better modularity.
2024–2025 — Mosaic GPU and ROCm infrastructure expansion
41 changes.
This period focused on establishing the Mosaic GPU backend as a first-class compilation path for NVIDIA hardware, introducing a dedicated MLIR dialect, MGPU kernels, and comprehensive build and CI infrastructure. Simultaneously, the project overhauled its ROCm support by migrating to a modular plugin architecture and manylinux-based wheel builds, while expanding TPU capabilities with new sparse and paged attention kernels.
2026 — OneAPI support and Pallas enhancements
15 changes.
This period focused on expanding hardware support with the initial implementation of the Intel OneAPI GPU backend and plugin infrastructure. Significant work was also done to enhance the Pallas compiler, including a new GPU kernel interpreter for debugging, experimental Triton lowerings for ragged operations, and precision testing. The underlying build system and CI pipelines were modernized to support Python 3.15, Bzlmod, and centralized test result analysis.
Features
Add C++ example for running JAX programs via the public XLA:CPU PJRT API
A new example in examples/jax\_cpp demonstrates how to load a compiled HLO module and execute it using the public XLA:CPU PJRT client from C++. The example shows how to use the new CompileAndLoad API, manage buffers with PjRtMemorySpace, and run computations on the CPU, providing a concrete reference for integrating JAX programs into C++ applications.
_examples/jax\cpp · high confidence
Add Cloud TPU Colab notebooks and landing page
This change introduces a new \cloud\_tpu\_colabs\ directory containing a landing page (README) and several example Jupyter notebooks for running JAX on Cloud TPUs. The included notebooks demonstrate core JAX capabilities such as \pmap\ for parallel execution (Pmap\_Cookbook), automatic differentiation and basic NumPy operations (JAX\_demo, JAX\_NeurIPS\_2020\_demo), and physics simulations like the Lorentz ODE Solver and Wave Equation. The README provides guidance on running these examples in Colab or on Cloud TPU VMs, including performance notes on padding and bfloat16 precision.
_cloud\_tpu\colabs · high confidence
Add Grouped Matrix Multiplication (GMM) kernels for TPU
Introduces new grouped matrix multiplication operations for TPU via the \jax.experimental.pallas.ops.tpu.megablox\ module. The \gmm\ function performs grouped matrix multiplications using Pallas kernels, with a custom VJP implementation in \ops.py\ that supports efficient backpropagation for both the left-hand side and right-hand side matrices (including transposed variants). The implementation includes device-specific logic in \common.py\ to handle TPU generation differences, such as bfloat16 support on TPU v4 and later, and exposes the functionality through the \megablox\ package namespace.
jax/experimental/pallas/ops/tpu/megablox · high confidence
Add JAX implementation of scipy.cluster.vq.vq
Users can now use jax.scipy.cluster.vq.vq to assign observation vectors to the nearest code in a code book based on Euclidean distance. This new function returns both the assigned code indices and the corresponding distances, enabling vectorized clustering operations within JAX programs.
_jax/\src/scipy/cluster · high confidence
Add Mosaic GPU example programs for matmul and Flash Attention
The \jax/experimental/mosaic/gpu/examples\ directory now contains runnable example programs demonstrating Mosaic GPU capabilities. This includes a pipelined matmul kernel (\matmul.py\) supporting H100 architectures, a Blackwell-specific matmul kernel (\matmul\_blackwell.py\) utilizing collective MMA and TMEM, and a Flash Attention implementation (\flash\_attention.py\) that supports Grouped Query Attention. These examples serve as reference implementations for using the Mosaic GPU compiler and runtime APIs.
jax/experimental/mosaic/gpu/examples · high confidence
Add Mosaic GPU matrix multiplication benchmarks
Added a new benchmark suite for Mosaic GPU matrix multiplication operations, including both bf16\_i8 and f32 precision variants. The \matmul\_bench.py\ script and its associated \BUILD\ file allow users to measure performance metrics such as TFLOPS and speedup for specific matrix dimensions and tiling configurations on H100 GPUs.
benchmarks/mosaic · high confidence
Add example for serving jax2tf models with TensorFlow Model Server
New files have been added to the jax2tf serving example directory, including a README with step-by-step instructions, a Python script (model\_server\request.py) for sending inference requests via gRPC or HTTP REST, and an \\init\\_.py module. This enables users to export JAX models converted to TensorFlow SavedModels and serve them using the open-source TensorFlow Model Server, supporting both fixed and batch-polymorphic models.
jax/experimental/jax2tf/examples/serving · high confidence
Add jax.scipy.spatial.transform Rotation class
JAX now provides a \Rotation\ class in \jax.scipy.spatial.transform\ that mirrors \scipy.spatial.transform.Rotation\. This new feature allows users to create, compose, and convert 3D rotations using quaternions, Euler angles, rotation matrices, and rotation vectors, with support for operations like inversion, application to vectors, and conversion between representations.
_jax/\src/scipy/spatial · high confidence
Added Pallas-MGPU planar\_snake benchmark
A new benchmark for the Pallas-Mosaic-GPU \planar\_snake\ calculation has been added to the \benchmarks/pallas\ directory. This includes a Python script (\mgpu\_planar\_snake\_bench.py\) that measures latency across various grid shapes and tile widths, along with the corresponding BUILD configuration to run the test on GPU H100 full configurations.
benchmarks/pallas · high confidence
Added Rotation and Slerp classes to scipy.spatial.transform
JAX now exposes the \Rotation\ and \Slerp\ classes within the \jax.scipy.spatial.transform\ module, allowing users to perform rotation operations and spherical linear interpolation in a JAX-compatible way.
jax/scipy/spatial · high confidence
Added auto-generated documentation for JAX converter evaluation results and primitive coverage
The jax2tf documentation now includes auto-generated reports on model conversion success rates across various backends (jax2tf\_xla, jax2tf\_noxla, jax2tfjs, jax2tflite, jax2tflite+flex) and detailed tables of JAX primitive support by data type. These files, along with their templates, are built from test harnesses to keep the limitations and coverage information up to date.
jax/experimental/jax2tf/g3doc · high confidence
Automatic distributed initialization for multiple cluster environments
JAX now automatically detects and initializes distributed runs in Slurm, Open MPI, MPI4py, Kubernetes, and Cloud TPU environments without requiring manual argument passing to \jax.distributed.initialize()\. A new \jax.\_src.clusters\ module provides a generic interface where specific cluster implementations (SlurmCluster, OmpiCluster, Mpi4pyCluster, K8sCluster, GkeTpuCluster, GceTpuCluster) are registered and checked in a defined order to auto-populate coordinator address, process count, and process ID. This allows users to launch distributed JAX jobs on these platforms with minimal configuration, as the system will handle the underlying discovery and connection setup automatically.
_jax/\src/clusters · high confidence
CI test results are now automatically loaded into BigQuery
A new post-processing pipeline in the \ci/postprocess\ directory now ingests JUnit XML test reports from Google Cloud Storage, converts them to JSON, and loads them into BigQuery tables (\jax\_ci.tests\, \jax\_ci.job\_metadata\, \jax\_ci.jobs\, and \jax\_ci.workflow\_runs\). This enables centralized storage and analysis of CI test outcomes, with automatic schema detection and support for both Bazel and Pytest report formats.
ci/postprocess · high confidence
Enable building JAXlib with Mosaic support
This change introduces the build infrastructure and C++ headers required to compile JAXlib with the Mosaic dialect. It adds a new \jaxlib/mosaic/BUILD\ file that defines C++ libraries for the TPU dialect, serialization passes, and C API bindings, while also exposing Python bindings via \jax.experimental.mosaic\. The diff establishes compatibility with the XLA Mosaic implementation by forwarding headers and build targets, allowing users to utilize Mosaic-specific optimizations and operations within JAXlib.
jaxlib/mosaic · high confidence
Experimental GPU lowering for ragged\_dot\_general via Pallas-Triton
Added an experimental Pallas-Triton implementation for the \ragged\_dot\_general\ operation on GPU. This new lowering, located in \jax/\_src/lax/pallas\_lowerings/gpu/ragged\_dot.py\, enables efficient computation of ragged matrix multiplications by dynamically handling variable-sized groups and chunking large groups across multiple streaming multiprocessors to improve hardware utilization.
_jax/\_src/lax/pallas\lowerings · high confidence
Experimental static key reuse checking
Added an experimental module \jax.experimental.key\_reuse\ that detects when a JAX random key is used more than once, raising a \KeyReuseError\ if detected. Users can enable this check globally via the \jax\_debug\_key\_reuse\ configuration flag or locally using the \jax.debug\_key\_reuse\ context manager. This feature helps prevent subtle bugs arising from incorrect PRNG key consumption patterns.
_jax/experimental/key\reuse · high confidence
Expose JAX-compatible scipy.optimize interface
The jax/scipy/optimize module now explicitly exports the minimize function and OptimizeResults class, re-exporting them from the internal implementation source. This change makes the optimization capabilities available under the public jax.scipy.optimize namespace, allowing users to import and use these functions directly from the standard scipy.optimize path while leveraging JAX's transformations.
jax/scipy/optimize · high confidence
Expose MLIR IR and PassManager modules in jaxlib
The \jax/\_src/lib/mlir\ package now explicitly re-exports the \ir\ and \passmanager\ modules from \jaxlib.mlir\, making these core MLIR components directly accessible to users and internal JAX code without needing to import from the deeper \jaxlib.mlir\ namespace.
_jax/\src/lib/mlir · high confidence
Expose Mosaic GPU MLIR dialect types and attributes via C API
The Mosaic GPU MLIR dialect is now accessible from external C/C++ code through a new C API. This exposes C-level functions to create, inspect, and manipulate key dialect components, including \TileTransformAttr\, \SwizzleTransformAttr\, various layout attributes (\WGSplatFragLayoutAttr\, \WGStridedFragLayoutAttr\, \TiledLayoutAttr\), and new data types such as \BarrierType\, \B6x16P32Type\, and \P2B6Type\. This enables other parts of the JAX ecosystem to interact directly with Mosaic GPU's internal representation.
jaxlib/mosaic/dialect/gpu/integrations · high confidence
Expose TPU dialect C API headers in jaxlib
The jaxlib build now includes the TPU dialect C API header (tpu\_dialect.h), which forwards declarations from the XLA Mosaic dialect. This enables jaxlib to link against and expose the C API functions for VectorLayout, VRegDataBounds, and the apply-vector-layout operations (assemble, disassemble, relayout, applyLayoutOp) to Python bindings in jax.experimental.mosaic.
jaxlib/mosaic/dialect/tpu/integrations · high confidence
Expose TPU dialect headers in jaxlib
The jaxlib/mosaic/dialect/tpu directory now exposes the TPU dialect definitions (including tpu.td, tpu\_ops.td, tpu\_types.td, tpu\_enums.td, and utility headers) by forwarding to the corresponding files in xla/mosaic/dialect/tpu. This allows jaxlib consumers to access the TPU dialect interface without directly depending on the XLA source tree.
jaxlib/mosaic/dialect/tpu · high confidence
Expose segment aggregation functions in jax.ops
Users can now import segment\_sum, segment\_prod, segment\_min, and segment\_max directly from jax.ops. These functions are re-exported from the internal jax.\_src.ops.scatter module to provide a public API for NumPy-style indexed segment aggregations.
jax/ops · high confidence
Expose sparse linear algebra solvers in jax.scipy.sparse.linalg
Users can now import and use the Conjugate Gradient (cg), Generalized Minimal Residual (gmres), and Biconjugate Gradient Stabilized (bicgstab) solvers directly from the public \jax.scipy.sparse.linalg\ module. This change establishes the public API surface for these sparse linear algebra operations, which are implemented in the internal source and re-exported for external consumption.
jax/scipy/sparse · high confidence
Initial OneAPI GPU backend support
This change introduces the foundational infrastructure for the Intel OneAPI (SYCL) GPU backend in JAX. It adds the build configuration (BUILD.bazel) and core runtime wrappers (oneapi\_gpu\_runtime) that map SYCL queue operations like memcpy, memset, and stream synchronization to JAX's internal status handling. The update also includes the plugin extension glue code (oneapi\_plugin\_extension) to register the backend with JAX's Python client, and implements the initial set of solver FFI kernels (oneapi\_solver\_kernels\_ffi) for linear algebra operations such as LU decomposition (getrf), enabling these computations to run on OneAPI-compatible hardware.
jaxlib/oneapi · high confidence
Introduce GPU kernel interpreter for Pallas
Adds a new GPU interpret mode that allows Pallas GPU kernels to be executed on the CPU for debugging and testing purposes. This feature simulates GPU-specific behaviors including shared memory spaces (GMEM, SMEM, TMEM), thread hierarchies (blocks, clusters, warps), and synchronization primitives (barriers). It also includes race detection for memory accesses and support for complex operations like TMA transfers and collective allocations, accessible via the \InterpretGPUParams\ configuration.
_jax/\_src/pallas/mosaic\gpu/interpret · high confidence
Introduce JAX Debugger with CLI, Colab, and Web backends
JAX now includes a new interactive debugger accessible via \jax.debug.breakpoint()\. This feature supports three backends: a command-line interface (CLI) debugger named \jdb\, a Colab-specific debugger for Jupyter environments, and a web-based debugger using \web\_pdb\. The debugger allows users to inspect stack frames, evaluate expressions, and view source code context. Users can select a specific backend via the \backend\ argument or let the system automatically choose the highest-priority available debugger. The implementation includes a registry system for debugger backends and handles frame filtering and token-based ordering.
_jax/\src/debugger · high confidence
Introduce JAX Mosaic experimental module with TPU custom call bindings
A new \jax.experimental.mosaic\ package has been added to expose bindings for TPU custom calls and MLIR dialects. This module re-exports key components from \jax.\_src.tpu\_custom\_call\, including \Tiling\, \OptLevel\, \TpuMemorySpace\, and \register\_extra\_dialect\, while also providing access to the TPU dialect via \jax.\_src.lib.tpu\. This change makes these lower-level TPU-specific APIs available for experimental use in Mosaic.
jax/experimental/mosaic · high confidence
Introduce JAX OneAPI plugin for Intel GPU support
This change adds the initial infrastructure for a OneAPI PJRT plugin, enabling JAX to run on Intel GPUs. The \jax\_plugins/oneapi\ directory now contains the build rules (\BUILD.bazel\) to compile the \pjrt\_c\_api\_gpu\_plugin.so\ native library, along with Python packaging scripts (\plugin\setup.py\, \setup.py\) that define the \jax-oneapi-plugin\ and \jax-oneapi-pjrt\ packages. The plugin's \\\init\\_.py\ implements logic to dynamically load required OneAPI runtime libraries (such as SYCL, MKL, and Level Zero adapters) and registers the plugin via JAX's plugin entry points, allowing users to utilize Intel hardware acceleration.
_jax\plugins/oneapi · high confidence
Introduce Mosaic GPU MLIR dialect for GPU kernel programming
Adds the \mosaic\_gpu\ MLIR dialect to JAX, providing a specialized set of operations and types for programming NVIDIA GPUs. This includes definitions for warp-group semantics, tensor memory (TMEM) management, asynchronous memory transfers (TMA), and collective matrix multiply (WGMMA) operations. The dialect enables lower-level control over GPU hardware features, supporting complex memory layouts and synchronization primitives required for high-performance custom kernels.
jaxlib/mosaic/dialect/gpu · high confidence
Introduce Mosaic GPU backend for JAX
Adds the Mosaic GPU backend to JAX, providing a new compilation and execution path for GPU kernels. This includes the core custom call implementation, MLIR pass pipelines for lowering GPU operations to PTX and SASS, serialization support for kernel caching, and debugging utilities for dumping intermediate compilation artifacts. The backend integrates with the XLA FFI and supports features like collective metadata and TMA descriptors.
jaxlib/mosaic/gpu · high confidence
Introduce MosaicGPU wheel build infrastructure
Added the build files and packaging scripts necessary to create a standalone MosaicGPU Python wheel. This includes a Bazel build rule to compile the \mosaic\_gpu.so\ shared library, a linker script to expose specific symbols, and a \setup.py\ configuration that defines the package structure, CUDA-specific dependencies, and an entry point for the plugin.
jaxlib/mosaic/gpu/wheel · high confidence
Introduce Pallas Mosaic GPU backend
Adds a new Pallas backend for NVIDIA GPUs via the Mosaic GPU library, providing the core abstractions, lowering logic, and build infrastructure needed to execute Pallas kernels on hardware. This includes the \jax.\_src.pallas.mosaic\_gpu\ package with modules for core types, lowering rules, primitives, and pipelining, enabling users to write and run GPU-accelerated custom kernels using the Mosaic GPU dialect.
_jax/\_src/pallas/mosaic\gpu · high confidence
Introduce Pallas as a first-party JAX module
Pallas is now included directly in the JAX source tree under \jax.\_src.pallas\, providing a unified location for the Pallas programming model. This change introduces the core Pallas primitives, the HLO interpreter for emulation, a cost estimation tool, and the \einshape\ primitive, alongside the necessary build rules to integrate these components into the JAX codebase.
_jax/\src/pallas · high confidence
Introduce Python bindings for the MLIR Triton dialect
This change adds the \jaxlib.triton\ module, providing Python bindings for the MLIR Triton dialect. It exposes core types like \PointerType\ and operations such as \ReduceOp\ and \ScanOp\ with improved return type inference, allowing users to interact with the Triton dialect directly within the JAX MLIR infrastructure.
jaxlib/triton · high confidence
Introduce Python bindings for the Mosaic TPU and GPU MLIR dialects
This change adds the initial Python bindings for the Mosaic TPU and GPU MLIR dialects, exposing them via \jax.experimental.mosaic\. The \BUILD\ file defines the build targets for \tpu\_dialect\ and \gpu\_dialect\, which generate Python wrappers for operations and enums. The \tpu.py\ module provides Python interfaces for TPU-specific operations like \vector\_load\, \vector\_store\, and \reinterpret\_cast\, while \mosaic\_gpu.py\ exposes the GPU dialect operations, including a custom \WarpMapOp\. Supporting files like \layout\_defs.py\ define necessary types such as \Direction\ and \ImplicitDim\ used by these bindings.
jaxlib/mosaic/python · high confidence
Introduce Splash Attention, a sparse flash attention kernel for JAX Pallas on TPU
Adds the \splash\_attention\ module to JAX Pallas, providing a general-purpose sparse flash attention kernel optimized for TPUs. This new capability allows users to specify attention masks using NumPy arrays and supports features such as dynamic masks, attention sinks (including backward pass support), and non-128 attention head dimensions (via padding). The module exposes key components including mask types (Causal, Local, Full, Random, ChunkedCausal), mask utilities, and kernel factories for MHA and MQA variants, enabling efficient attention computation with flexible masking strategies.
_jax/experimental/pallas/ops/tpu/splash\attention · high confidence
Introduce TPU Ragged Paged Attention kernel with auto-tuned block sizes
Adds a new \ragged\_paged\_attention\ kernel in \jax.experimental.pallas.ops.tpu\ designed for high-throughput inference on TPUs, supporting mixed prefill and decoding workloads. The implementation includes a reference implementation for validation, input validation logic, and a comprehensive auto-tuned block size table (\tuned\_block\_sizes.py\) optimized for TPU v6 configurations to maximize performance and avoid register spilling.
_jax/experimental/pallas/ops/tpu/ragged\_paged\attention · high confidence
Introduce dedicated ROCm build infrastructure and plugin extension
This change establishes the \jaxlib/rocm\ directory with a new \BUILD\ file defining HIP-based GPU kernels (such as RNN, solver, and linear algebra operations) and a \rocm\_plugin\_extension.cc\ that registers the ROCm-specific FFI handlers and device utilities. It also adds \rocm\_rpath.bzl\ to configure wheel-relative RUNPATHs for AMD ROCm libraries, ensuring the JAX wheel can locate dependencies at runtime, and \rocm\_version.bzl\ to expose the detected ROCm version to the build system.
jaxlib/rocm · high confidence
Introduce experimental Colocated Python API
Adds a new \jax.experimental.colocated\_python\ module that enables executing Python functions and classes on the same devices as JAX arguments. The API exposes \colocated\_cpu\_devices\ to locate CPU devices colocated with accelerators, a \colocated\_python\ decorator to wrap serializable Python functions for execution on those devices, and a \colocated\_python\_class\ decorator to wrap Python classes so their methods run on the backend. The implementation includes serialization logic for handling JAX arrays, meshes, and devices, along with backend components for managing object lifecycles and results.
_jax/experimental/colocated\python · high confidence
Introduce experimental JAX roofline API for performance modeling
JAX now includes an experimental roofline analysis API in \jax.experimental.roofline\ that allows users to estimate the computational and memory performance characteristics of their JAX programs. This new module provides functions like \roofline()\ and \roofline\_and\_grad()\ to calculate metrics such as FLOPs, HBM (High Bandwidth Memory) bytes, and inter-chip interconnect (ICI) latency by interpreting the JAXPR. It supports a wide range of primitives including unary and binary operations, convolutions, gather/scatter, and various reduction operations, enabling developers to profile and optimize their models for specific hardware constraints.
jax/experimental/roofline · high confidence
Introduce experimental Pallas Fuser API for manual kernel fusion
This change adds a new \jax.\_src.pallas.fuser\ module that provides an experimental API for manually fusing computations into Pallas kernels. It introduces the \@fuser.fuse\ decorator to fuse surrounding computation into a \@fuser.fusible\ block, and the \@fuser.custom\_fusion\ decorator to define custom fusion behavior for specific operations. The module also includes internal utilities for block spec propagation, JAXPR fusion, and custom evaluation, enabling users to optimize kernel performance by controlling fusion granularity.
_jax/\src/pallas/fuser · high confidence
Introduce experimental array and pytree serialization library
This change introduces a new \jax.experimental.array\_serialization\ module that provides asynchronous checkpointing capabilities for JAX arrays and nested pytrees. It exposes a \GlobalAsyncCheckpointManager\ for managing concurrent serialization and deserialization operations, along with dedicated \save\ and \load\ functions for pytrees. The implementation leverages TensorStore with Zarr3 and zstd compression to handle storage, and includes utilities for serializing pytree structures and managing memory limits during I/O operations.
_jax/experimental/array\serialization · high confidence
Introduce experimental sparse matrix support in JAX
The new \jax.experimental.sparse\ submodule provides experimental support for sparse matrix operations, centered on the \BCOO\ (batched coordinate) sparse array type and the \sparsify\ transform. Users can now create sparse arrays from dense data, perform sparse-dense and sparse-sparse matrix products, and apply JAX transformations like \jit\, \vmap\, and \grad\ directly to sparse objects. The \sparsify\ transform allows existing dense JAX functions to operate on sparse inputs by automatically routing supported primitives (such as matrix multiplication, elementwise operations, and reductions) to their sparse equivalents. Additionally, the module includes sparse-aware versions of \jax.grad\, \jax.value\_and\_grad\, \jax.jacfwd\, and \jax.jacrev\ that compute gradients within the subspace defined by the array's sparsity pattern, and provides low-level GPU lowerings via cuSPARSE/hipSPARSE for efficient sparse matrix-vector and matrix-matrix products.
jax/experimental/sparse · high confidence
Introduce jax.extend module for internal JAX machinery access
The new \jax.extend\ package provides a structured set of submodules (backend, core, linear\_util, lowering, pallas, random, sharding, source\_info\_util, and xla) that expose internal JAX components for extension purposes. This module is explicitly documented as having no compatibility guarantee across releases, meaning users relying on these internals should expect breaking changes. The package includes specific utilities such as PRNG implementation definition, HLO module transformation registration, and sharding proto serialization helpers, all organized under a new Bazel build target.
jax/extend · high confidence
Introduce jax.image public API with resize and scale\_and\_translate
Users can now import image manipulation functions directly from the \jax.image\ namespace. This new module exposes \resize\ (including the \ResizeMethod\ enum) and \scale\_and\_translate\, providing a public interface for image scaling and translation operations that are implemented using native JAX primitives to support vmap and gradients.
jax/image · high confidence
Introduce jax2tf package with JAX-to-TensorFlow conversion and TensorFlow-to-JAX calling capabilities
The jax2tf package is introduced to provide bidirectional interoperation between JAX and TensorFlow. It adds the jax2tf.convert API, which allows JAX functions to be called within a TensorFlow context (eager or graph) or serialized as a TensorFlow SavedModel, and the jax2tf.call\_tf API, which enables calling TensorFlow functions from within JAX with support for reverse-mode autodiff. The package defaults to native serialization using StableHLO and the XlaCallModule op for high fidelity and performance, while also including build targets, documentation, and a deprecation notice for the legacy getting-started notebook.
jax/experimental/jax2tf · high confidence
Introduce mutable array state primitives and discharge machinery
JAX now supports mutable arrays via a new state system located in \jax/\_src/state\. This change introduces \Ref\ types (represented by \AbstractRef\ and \TransformedRef\) that allow in-place mutation of array data through primitives like \get\, \swap\, and \addupdate\. The system includes a state discharge mechanism (\discharge\_state\) that converts stateful JAX computations into pure ones by threading updates through the computation, enabling these mutable operations to work within JAX's functional transformation framework (e.g., \jit\, \grad\).
_jax/\src/state · high confidence
Introduction of the jax.export module for serializing JAX functions
The new \jax.\_src.export\ package provides APIs to export JAX functions into a portable, serialized format (StableHLO) for interoperation with other runtimes. This includes the \Exported\ dataclass to hold lowered modules, FlatBuffers-based serialization (\serialization.fbs\) with version 12 support, and shape polymorphism handling (\shape\_poly.py\) to allow functions with symbolic dimensions. Users can now serialize JAX programs with explicit sharding, memory spaces, and PRNG keys, and deserialize them for execution on different platforms, with backward compatibility guarantees for the serialization schema.
_jax/\src/export · high confidence
New Bazel repository rule for external test dependencies
A new Bazel repository rule, \external\_deps\_repository\, has been added to \third\_party/external\_deps\ to streamline the configuration of external test dependencies. This rule allows users to pass a list of dependency targets (e.g., ROCm PJRT imports) and automatically generates a \.bzl\ file containing an \EXTERNAL\_DEPS\ variable, which can then be loaded in other BUILD files to access those targets. This change supports the project's shift toward bzlmod by providing a structured way to manage external dependencies within the module system.
_third\_party/external\deps · high confidence
New Bazel-based wheel build infrastructure for JAX and GPU plugins
The \jaxlib/tools\ directory now contains the build scripts and Bazel rules that assemble the \jaxlib\, CUDA, ROCm, OneAPI, and Mosaic GPU wheels. This change introduces a new build system using \jax\_wheel\ and \wheel\_sources\ Bazel targets, replacing previous ad-hoc methods. It adds dedicated build scripts (\build\_wheel.py\, \build\_gpu\_kernels\_wheel.py\, \build\_gpu\_plugin\_wheel.py\, \build\_mosaic\_wheel.py\) and utilities (\build\_utils.py\) to handle source collection, platform-specific configuration (CUDA/ROCm/OneAPI versions), and editable installs. It also includes a \wheel\_size\_test.py\ to enforce maximum wheel size limits and a \rocm\_wheel\_deps.bzl\ rule to aggregate ROCm runtime dependencies. This infrastructure ensures reproducible wheel content and filenames, supports hermetic Python builds, and centralizes the logic for packaging JAX's native extensions and plugins.
jaxlib/tools · high confidence
New CI infrastructure and documentation for JAX builds and tests
The \ci/\ directory now contains a comprehensive set of shell scripts and documentation that standardize how JAX artifacts are built and tested. New scripts handle building wheels for CUDA, ROCm, and OneAPI backends, while dedicated runners execute Bazel tests on CPU and GPU targets using both Remote Build Execution (RBE) and local strategies. The system is supported by new documentation (\CONTRIBUTING.md\, \README.md\) explaining the hybrid CI architecture, and utility scripts for parsing wheel metadata and configuring environment variables.
ci · high confidence
New CI utility scripts for build, testing, and environment management
This change introduces a new set of utility scripts in the \ci/utilities\ directory to support JAX's continuous integration workflows. Key additions include \setup\_build\_environment.sh\ for initializing the build environment and cloning XLA, \install\_wheels\_locally.sh\ for installing JAX wheels using \uv\, and \run\_docker\_container.sh\ for managing Docker containers with environment variable passing and Windows MSYS path conversion. Testing support is enhanced with \collect\_bazel\_test\_xmls.sh\ to normalize Bazel test results, \prepare\_rocm\_tests.sh\ and \rocm\_test\_env.sh\ for ROCm-specific test setup, and \setup\_portserver.sh\ to manage the portserver daemon. Additional utilities include \run\_auditwheel.sh\ for manylinux compliance checks, \set\_artifact\_tag\_flags.sh\ for wheel version stamping, \generate\_invocation\_id.py\ for Bazel invocation IDs, \report\_resultstore\_link.py\ for ResultStore reporting, and \convert\_msys\_paths\_to\_win\_paths.py\ for Windows path handling. A README.md documents these scripts.
ci/utilities · high confidence
New JAX implementations for scipy.interpolate, scipy.linalg, and scipy.special
This change adds JAX-compatible implementations for several SciPy modules located in \jax/\_src/third\_party/scipy\. It introduces \RegularGridInterpolator\ for interpolating points on a regular rectangular grid (supporting linear and nearest methods), \funm\ for evaluating matrix-valued functions via Schur decomposition, and \betaln\ for computing the log of the beta function with improved accuracy for large inputs. Additionally, it provides a real-valued implementation of the Fresnel integrals (\fresnel\) in \scipy.special\, helper utilities for spectral analysis in \signal\_helper\, and the necessary license and initialization files to support these new third-party components.
_jax/\_src/third\party/scipy · high confidence
New JAX implementations for scipy.stats distributions and utilities
This change introduces JAX-compatible implementations for a broad set of scipy.stats functions, including statistical distributions (Bernoulli, Beta, Binomial, Cauchy, Chi-square, Dirichlet, Exponential, Gamma, Generalized Normal, Geometric, and Beta-Binomial) and utility functions (mode and rankdata). Users can now compute probability mass/density functions (pmf/pdf), cumulative distribution functions (cdf), survival functions (sf), and their logarithmic variants for these distributions using JAX arrays, enabling differentiable statistical modeling and JIT compilation.
_jax/\src/scipy/stats · high confidence
New JAX tools for sampler bias evaluation and IR conversion
This change introduces new utility scripts in the jax/tools directory. The evaluate\_sampler\_bias.py tool allows users to evaluate the accuracy of JAX random samplers by comparing empirical quantiles against analytical expectations from scipy, supporting various distributions. Additionally, jax\_to\_ir.py provides a way to convert JAX functions into serialized IR formats (HLO or TensorFlow graphs) for use in ahead-of-time compilation or external systems. A new pgo\_nsys\_converter.py script is also added to convert NVIDIA Nsight Systems profiles into the .pbtxt format for XLA's Profile Guided Latency Estimator.
jax/tools · high confidence
New L-BFGS optimizer and stricter minimize API
The \jax.scipy.optimize.minimize\ function now supports the L-BFGS algorithm via the method name \l-bfgs-experimental-do-not-rely-on-this\, in addition to the existing BFGS solver. This new optimizer is implemented in \jax/\_src/scipy/optimize/\_lbfgs.py\ and exposes detailed iteration history (such as step and gradient histories) in its results. Additionally, the \minimize\ function now strictly enforces that the \args\ argument must be a tuple, raising a \TypeError\ if a list or other sequence is provided.
_jax/\src/scipy/optimize · high confidence
New MLIR Python bindings and Mosaic GPU dialect support
This change introduces a new build target and set of Python extension modules in \jaxlib/mlir/\_mlir\_libs\ that expose MLIR dialects and the new Mosaic GPU compiler infrastructure to Python. It adds nanobind-based bindings for the core MLIR library, GPU/NVGPU/LLVM/SparseTensor dialects, and specific JAX extensions for TPU (\\_tpu\_ext\), Triton (\\_triton\_ext\), and Mosaic GPU (\\_mosaic\_gpu\_ext\). These modules provide Python APIs for registering dialects, handling MLIR types and attributes (such as Mosaic's \BarrierType\, \TileTransformAttr\, and \SwizzleTransformAttr\), and managing traceback-to-location mapping for better error reporting. The build system also includes automated generation of \.pyi\ type stubs for these extensions to improve type checking for users.
_jaxlib/mlir/\_mlir\libs · high confidence
New Mosaic GPU kernel implementations for multi-GPU operations
This change introduces a new \jax/experimental/pallas/ops/gpu\ package containing high-performance, multi-GPU (MGPU) kernels built on the Mosaic GPU backend. Users can now leverage specialized implementations for collective matrix multiplication (\collective\_matmul\_mgpu\), all-gather operations (\all\_gather\_mgpu\), and FlashAttention-3 (\attention\_mgpu\) that utilize hardware-specific features like TMA (Tensor Memory Accelerator) and warpgroup semantics for improved performance on Hopper and Blackwell architectures. The module also includes ragged dot kernels (\blackwell\_ragged\_dot\_mgpu\) for grouped matrix multiplications, providing a unified location for these advanced MGPU capabilities within the Pallas ecosystem.
jax/experimental/pallas/ops/gpu · high confidence
New Pallas TPU Interpret Mode for CPU-based debugging
Adds a new TPU interpret mode that allows Pallas TPU kernels to be executed on the CPU for debugging and testing purposes. This feature introduces a new \jax/\_src/pallas/mosaic/interpret\ module containing the interpreter logic, shared memory simulation, race detection, and synchronization primitives (semaphores, barriers). Users can enable this mode by passing an \InterpretParams\ instance to \pallas\_call\ or \core\_map\, which simulates TPU-specific behaviors like HBM/VMEM memory spaces, remote/local DMAs, and vector-clock-based race detection, while providing configuration options for logging, out-of-bounds handling, and floating-point operation skipping.
_jax/\src/pallas/mosaic/interpret · high confidence
New Pallas TPU kernels for Flash Attention, All-Gather, and Matmul
This change introduces new experimental Pallas kernels for TPUs, including a Flash Attention implementation (supporting causal masking, segment IDs, and head dimensions up to 128), a pedagogical All-Gather collective kernel, and a basic Matmul kernel. These are provided as example implementations in \jax.experimental.pallas.ops.tpu\ to demonstrate how to write custom Pallas operations using \pl.pallas\_call\, \pltpu\ primitives, and \shard\_map\.
jax/experimental/pallas/ops/tpu · high confidence
New Pallas-based PRNG implementations for TPU
Added new Pallas kernel implementations for the Philox and Threefry pseudo-random number generators in \jax.experimental.pallas.ops.tpu.random\. These modules provide hardware-accelerated random bit generation on TPUs, exposing the \pallas\_threefry2x32\ PRNG implementation via the standard JAX PRNG interface. The change includes the kernel logic, utility functions for block-based indexing, and registration of the new Threefry implementation, enabling users to leverage Pallas for high-performance random number generation on TPU hardware.
jax/experimental/pallas/ops/tpu/random · high confidence
New ROCm build and wheel-fixing tooling
Added a suite of Python and shell scripts in build/rocm/tools to streamline ROCm wheel production and validation. build\_wheels.py orchestrates jaxlib and jax wheel builds (supporting both GCC and Clang compilers) via build.py, while fixwheel.py repairs wheels to be manylinux-compatible using auditwheel. Supporting utilities include get\_rocm.py for installing ROCm on Ubuntu and RHEL-based systems, blacken.sh for code formatting, libc.py and symbols.py for detecting glibc versions and symbols, and auditwheel integration to ensure binary compatibility.
build/rocm/tools · high confidence
New ROCm build infrastructure and CI tooling
This change introduces a comprehensive set of build scripts, Dockerfiles, and configuration files for building and testing JAX with ROCm support. It adds Dockerfiles for Ubuntu 22.04 and 24.04, a manylinux wheel build setup, and CI scripts for single and multi-GPU testing. The infrastructure also includes Bazel configuration for ROCm builds, test target definitions, and documentation for building JAX from source with ROCm.
build/rocm · high confidence
New TPU-specific linear algebra implementations for SVD, eigendecomposition, and QDWH
JAX now includes a dedicated \jax.\_src.tpu.linalg\ module containing JIT-compatible implementations of Singular Value Decomposition (SVD), symmetric eigendecomposition (eigh), and QR-based Dynamically Weighted Halley (QDWH) polar decomposition. These algorithms are optimized for TPU hardware, utilizing iterative methods like QR and Cholesky decompositions to avoid matrix inversion, and include support for computing subsets of singular values via \subset\_by\_index\. This change introduces the underlying computational primitives for these linear algebra operations on TPUs, distinct from the general CPU/GPU implementations.
_jax/\src/tpu · high confidence
New and updated example scripts for variational inference, differential privacy, and SPMD training
The examples directory now includes new scripts for Automatic Differentiation Variational Inference (advi.py), Differentially Private SGD (differentially\_private\_sgd.py), and SPMD MNIST classification (spmd\_mnist\_classifier\_fromscratch.py), alongside an ONNX-to-XLA compiler demo (onnx2xla.py). Existing examples have been updated to use the new \jax.example\_libraries\ namespace (replacing \jax.experimental\ and \jax.api\), migrated to \jax.numpy\ (jnp), and modernized to use \jax.random.key\ and \random.fold\in\ for PRNG management. The MNIST VAE example now uses the \optimizers\ library instead of \minmax\, and the MNIST classifier example has increased its default batch size to 128 and removed the \absl\ dependency in favor of standard \if \\name\\_ == '\_\main\\_'\ execution.
examples · high confidence
New cuDNN integration module for fused attention and scaled matmul
JAX introduces a new \jax.\_src.cudnn\ package that provides low-level lowering infrastructure for cuDNN operations. This includes a \cudnn\_fusion\ decorator to lower computations to XLA cuDNN fusions, a \fused\_attention\_stablehlo\ module implementing the scaled dot-product attention (SDPA) API with support for multiple input layouts (BTNH, BNTH) and mask types, and a \scaled\_matmul\_stablehlo\ module enabling block-scaled matrix multiplication with specific lowerings for CUDA, ROCm, and CPU platforms.
_jax/\src/cudnn · high confidence
New end-to-end JAX FFI example project with CPU, CUDA, and stateful demos
A new \examples/ffi\ directory provides a complete, buildable project demonstrating JAX's Foreign Function Interface. It includes C++ and CUDA implementations for several use cases: a basic \rms\_norm\ operation with custom VJP support, CPU examples showing global state caching (\counter\), attribute passing (\array\_attr\, \dictionary\_attr\), and input-output aliasing, as well as a CUDA example (\foo\_fwd\/\foo\_bwd\) and a GPU stateful example (\state\). The project uses CMake and nanobind to build the native extensions and registers them via \jax.ffi\, offering a practical reference for packaging and testing FFI extensions.
examples/ffi · high confidence
New example\_libraries module for Stax and optimizers
JAX now includes a dedicated \jax.example\_libraries\ package containing \stax\ and \optimizers\ modules, moved from \jax.experimental\. These mini-libraries serve as lightweight, educational examples for building neural networks and implementing first-order optimizers using JAX's functional API, rather than production-grade tools.
_jax/example\libraries · high confidence
New internal test utilities for JAX primitives and export compatibility
The \jax.\_src.internal\_test\_util\ package has been introduced to centralize internal testing infrastructure. It includes \test\_harnesses.py\, which defines a \Harness\ class for specifying inputs and callables to exercise JAX primitives, and \lax\_test\_util.py\, which provides LAX-specific test utilities. Additionally, \export\_back\_compat\_test\_util.py\ adds tools for verifying the backward compatibility of JAX serialized formats and custom calls, ensuring that changes do not break previously exported models.
_jax/\_src/internal\_test\util · high confidence
New jax.extend.core API for primitives and trace introspection
JAX now exposes a new \jax.extend.core\ module that provides public access to core primitives (such as \jit\_p\, \call\_p\, and various \lax\ primitives like \polynomial\_p\ and \log2\_p\) and trace introspection utilities (including \TraceTag\, \set\_current\_trace\, and \take\_current\_trace\). This change consolidates initial and final style custom VJP primitives and re-exports RNG primitives, enabling advanced extension and debugging capabilities that were previously internal.
jax/extend/core · high confidence
New jax.extend.mlir module for public MLIR access
A new \jax.extend.mlir\ package has been introduced to provide a stable, public-facing interface for JAX's MLIR capabilities. This module re-exports core MLIR components—such as the IR and pass manager—from \jaxlib\—alongside key JAX-specific functions like \lower\_with\_sharding\_in\_types\, \serialize\_portable\_artifact\, and \hlo\_to\_stablehlo\. This change establishes a dedicated extension point for users to interact with MLIR internals without relying on internal implementation details.
jax/extend/mlir · high confidence
New jax.extend.mlir.dialects package exposes MLIR dialects
The \jax.extend.mlir.dialects\ module has been added, providing a public-facing Python package that re-exports MLIR dialect bindings (arith, builtin, chlo, func, math, memref, mpmd, scf, sdy, sparse\_tensor, stablehlo, and vector) from \jaxlib.mlir.dialects\. This change establishes a new import path for these dialects, replacing the previous reliance on internal implementation paths.
jax/extend/mlir/dialects · high confidence
New jax.scipy.stats module with distribution exports
The jax/scipy/stats directory now contains a public API that exposes a wide range of statistical distributions (including bernoulli, beta, binom, cauchy, chi2, dirichlet, expon, gamma, gennorm, geom, gumbel\_l, gumbel\_r, laplace, logistic, multinomial, multivariate\_normal, nbinom, norm, pareto, poisson, t, truncnorm, uniform, vonmises, and wrapcauchy) as well as utility functions like gaussian\_kde, mode, rankdata, and sem. Each distribution module re-exports specific methods (such as pdf, cdf, ppf, logpdf, etc.) from the internal implementation, making them directly accessible via jax.scipy.stats.
jax/scipy/stats · high confidence
New jax2tf examples for SavedModel, Keras reuse, and TensorFlow.js
The \jax/experimental/jax2tf/examples\ directory now includes a comprehensive set of examples demonstrating how to save JAX models as TensorFlow SavedModels, reuse them in larger Keras pipelines, and convert them to TensorFlow.js. The new \saved\_model\_lib.py\ provides a \convert\_and\_save\_model\ helper that allows users to save model parameters as separate variables (enabling fine-tuning and avoiding GraphDef size limits) and supports shape polymorphism. The \keras\_reuse\_main.py\ example shows how to load a jax2tf SavedModel as a \hub.KerasLayer\ and train a classifier on top of it in TensorFlow. Additionally, a Quickdraw CNN example demonstrates conversion to TensorFlow.js, and the MNIST examples (\mnist\_lib.py\, \saved\_model\_main.py\) have been updated to use \optax\ instead of \flax.optim\ and support both pure JAX and Flax model implementations.
jax/experimental/jax2tf/examples · high confidence
New multi-pass source mapper for Jaxprs and HLO
The \jax.experimental.source\_mapper\ module introduces a new framework for generating source maps across multiple compiler passes. It provides a registry of passes (e.g., \jaxpr\, \hlo:stable-hlo\, \hlo:original\) that can be executed to produce source maps for JAX intermediate representations. The module includes specific generators for Jaxprs and HLO dialects, handling both old and new HLO metadata formats, and exposes utilities to compile functions with specific environment flags and generate dump outputs containing the source map and generated code.
_jax/experimental/source\mapper · high confidence
New suite of microbenchmarks for JAX core operations
Added a comprehensive set of new benchmark scripts to the \benchmarks/\ directory to measure the performance of various JAX capabilities. These include \api\_benchmark.py\ for JIT and eager dispatch, \math\_benchmark.py\ for unary and binary math operations, \linalg\_benchmark.py\ for linear algebra functions (SVD, QR, LU, etc.), \random\_benchmark.py\ for PRNG key handling, \sparse\_benchmark.py\ for BCOO/BCSR sparse operations, \shape\_poly\_benchmark.py\ for symbolic shape manipulation, \cache\_benchmark.py\ for internal caching, and \tracing\_benchmark.py\ for tracing and lowering overheads. All benchmarks utilize the \google\_benchmark\ library and include proper device-skip logic and \block\_until\_ready\ calls.
benchmarks · high confidence
New type stubs for LAPACK eigenvalue, Schur, and SVD modules
Added Python type stubs (\.pyi\) for the \jaxlib.cpu.\_lapack\ package, defining the \eig\, \schur\, and \svd\ submodules. These stubs introduce \ComputationMode\ and \Sort\ enums that specify options for computing eigenvectors, Schur vectors, and singular value matrices (full, min, or none), providing static type information for these linear algebra operations.
_jaxlib/cpu/\lapack · high confidence
Open-source PagedAttention TPU kernel with quantization and soft capping support
The PagedAttention TPU kernel is now available as an open-source component in JAX, providing a high-performance attention implementation optimized for TPU hardware. This release introduces support for quantized K/V pages (int8), allowing users to leverage reduced memory footprint and potentially faster inference with quantized models. Additionally, the kernel now supports attention logits soft capping, enabling better control over attention distribution for specific model architectures. The implementation includes a reference grouped query attention (GQA) implementation for validation and testing purposes.
_jax/experimental/pallas/ops/tpu/paged\attention · high confidence
jax.scipy restructuring and API expansion
The jax.scipy package has been restructured to expose a significantly expanded set of submodules and functions, including new modules for cluster (vq), fft (dct, idct), integrate (trapezoid), and interpolate (RegularGridInterpolator), alongside a comprehensive linalg module with support for decompositions (LU, QR, Cholesky, SVD, Schur, expm), special matrices, and solvers. The public API now imports these from internal sources, while deprecated or removed legacy modules like misc and stats have been deleted, and specific functions such as lpmn and lpmn\_values are now explicitly deprecated with warnings.
jax/scipy · high confidence
Removals
Deprecate jax.lib.xla\_bridge and migrate to internal source
The public module jax.lib.xla\_bridge has been removed, and its functionality has been moved to the internal package jax.\_src.lib. Users relying on jax.lib.xla\_bridge will encounter import errors; they should migrate to the new internal paths or use the stable public APIs provided by jax directly. The jax.lib package now re-exports the version string from the internal source.
jax/lib · high confidence
Architecture
Consolidate experimental build targets into jax/experimental/BUILD
The build configuration for experimental modules has been consolidated into a single \jax/experimental/BUILD\ file. This change introduces visibility controls via package groups (e.g., \buffer\_callback\_users\, \mosaic\_users\) and re-exports internal implementation targets (such as \//jax/\_src:buffer\_callback\) to the public experimental namespace, streamlining the build structure for experimental APIs.
jax/experimental · high confidence
JAX neural network module relocated to jax.\_src.nn
The implementation of the JAX neural network module (jax.nn), including activation functions, initializers, and attention utilities, has been moved from the public jax/nn directory into the internal source package jax/\_src.nn. This change establishes the new internal layout for these components, which are still accessible via the public jax.nn API.
_jax/\src/nn · high confidence
Move JAX ops implementation to jax.\_src.ops
The implementation of JAX's indexed update and segment operations (including scatter, gather, and segment reductions) has been moved from the public jax.ops module into the internal jax.\_src.ops package. This change consolidates the source code for these operations into a dedicated submodule, separating the internal implementation details from the public API surface while maintaining the existing functionality for users.
_jax/\src/ops · high confidence
Refactor control flow primitives into a dedicated module
The control flow primitives (scan, while\_loop, fori\_loop, cond, switch, custom\_root, custom\_linear\_solve, and cumulative reductions) have been reorganized from the main lax module into a new subpackage at jax/\_src/lax/control\_flow. This change splits the implementation into separate files for loops, conditionals, and solvers, and introduces a common utilities module for shared logic like tree/avals checking and constant merging. The public API surface remains unchanged, but the internal structure is now modularized to improve maintainability and reduce coupling within the lax package.
_jax/\_src/lax/control\flow · high confidence
Refactor lax module into dedicated submodules
The implementation of the \jax.lax\ module has been reorganized by splitting the monolithic \lax.py\ file into specialized submodules, including \jax.\_src.lax.convolution\ for convolution operations and \jax.\_src.lax.slicing\ for slicing, update\_slice, gather, and scatter operations. This structural change improves code maintainability and separation of concerns within the core lax primitives.
_jax/\src/lax · high confidence
Restructure jaxlib GPU build and kernel registration
The jaxlib/gpu directory has been reorganized to consolidate shared CUDA and ROCm GPU kernels into a new build target. This change introduces a dedicated BUILD file that defines the \gpu\_plugin\_extension\ library and exports a comprehensive set of kernel source files (including linear algebra, sparse, RNN, and Triton kernels). It also adds new helper infrastructure such as \handle\_pool\ for managing GPU library contexts, \ffi\_wrapper\ for legacy kernel compatibility, and \gpu\_kernel\_helpers\ for unified error handling across CUDA and ROCm backends. Additionally, GPU-specific FFI handlers are now explicitly registered in \gpu\_kernels.cc\, centralizing the dispatch logic for operations like LU decomposition, SVD, and sparse matrix-vector products.
jaxlib/gpu · high confidence
jax.numpy namespace restructured into modular subpackages
The \jax.numpy\ package has been refactored from a single monolithic \lax\_numpy.py\ module into a structured namespace with dedicated submodules. Implementation logic is now split into \jax.\_src.numpy.lax\_numpy\ (core operations), \jax.\_src.numpy.array\_constructors\ (e.g., \array\, \asarray\), \jax.\_src.numpy.array\_creation\ (e.g., \zeros\, \linspace\), \jax.\_src.numpy.einsum\, and \jax.\src.numpy.indexing\. Additionally, \jax.numpy.fft\ and \jax.numpy.linalg\ are now explicit subpackages that re-export their respective functions. A new \jax/numpy/\\init\\_.pyi\ type stub file has been added to provide static type checking for the public API.
jax/numpy · high confidence
jaxlib build system restructured with new Bazel rules and module organization
The jaxlib build configuration has been significantly reorganized. The BUILD file now loads a comprehensive set of custom macros from jaxlib:jax.bzl (such as nanobind\_extension, pytype\_strict\_library) and jaxlib:pywrap.bzl, replacing previous inline or external definitions. The package visibility is locked down to //jax:internal by default, and the main jaxlib pytype\_strict\_library now explicitly depends on a wide array of internal MLIR dialects (e.g., stablehlo\_dialect, chlo\_dialect, sdy\_dialect) and C++ extensions (e.g., \_jax, \_ifrt\_proxy, \_pathways, \_pretty\_printer). This reflects a shift towards a more modular, strictly typed, and internally encapsulated build structure for the jaxlib Python extension.
jaxlib · high confidence
Behavioural changes
CPU linear algebra and sparse kernels moved to XLA FFI with nanobind bindings
The CPU backend in jaxlib/cpu has been restructured to implement LAPACK, BLAS, tridiagonal solve, and CSR sparse kernels via the XLA Foreign Function Interface (FFI) instead of legacy custom calls. This change introduces new C++ kernel implementations (e.g., lapack\_kernels.cc, sparse\_kernels.cc) and Python extension modules (\_lapack, \_sparse) built with nanobind, enabling better integration with SciPy's LAPACK/BLAS libraries and supporting features like pivoted QR factorization and parallel batch processing for linear algebra operations on CPU devices.
jaxlib/cpu · high confidence
Deprecate jax.experimental.compilation\_cache in favor of jax.\_src.compilation\_cache
The public experimental compilation cache module (jax.experimental.compilation\_cache) now acts as a thin wrapper that re-exports only set\_cache\_dir and reset\_cache from the internal implementation (jax.\_src.compilation\_cache). This change signals that the experimental module is being deprecated and that users should migrate to the internal source for full functionality, while the experimental entry point remains available for backward compatibility with these two specific functions.
_jax/experimental/compilation\cache · high confidence
Deprecate jax.experimental.pjit APIs and prune exports
The \jax.experimental.pjit.NamedSharding\ and \jax.experimental.pjit.PartitionSpec\ classes are now deprecated in favor of their counterparts in \jax.sharding\. Additionally, several internal or experimental symbols previously exported from \jax.experimental.pjit\ have been removed from the public API surface.
_jax/\src · high confidence
Deprecation of the Pallas Triton GPU backend
The Pallas Triton backend for GPU execution is now deprecated and will be removed in a future JAX version. Users relying on \pl.pallas\_call\ with the Triton lowering path will see a deprecation warning at runtime, and are advised to migrate to the Mosaic GPU backend for Pallas or use the official Triton bindings via \jax\_triton\.
_jax/\src/pallas/triton · high confidence
Document and centralize JAXCI environment variables
The CI environment configuration is now centralized in \ci/envs/default.env\ and \ci/envs/docker.env\, with a new \ci/envs/README.md\ documenting all \JAXCI\\\ variables. This change introduces explicit support for OneAPI builds (defaulting to version 2025.1), allows configuring the Bazel output base to prevent disk space issues, and provides granular control over test execution via variables like \JAXCI\_IGNORE\_TESTS\ and \JAXCI\_EXTRA\_TEST\_ENV\. Docker-specific settings are also standardized, pointing to the new \ml-build\ container images.
ci/envs · high confidence
Establishes public API and type stubs for jax.nn
The \jax.nn\ module now exposes a comprehensive public interface for neural network operations, including activation functions (such as \silu\, \mish\, \squareplus\, and \log1mexp\), normalization utilities like \standardize\, and attention helpers like \dot\_product\attention\. This change introduces explicit type stubs (\\\init\\_.pyi\) to provide static type checking support and reorganizes the module structure by moving the core implementation into \jax.\_src.nn\ while maintaining backward-compatible exports. Users can now rely on consistent, typed access to these common NN functions directly from the \jax.nn\ namespace.
jax/nn · high confidence
Expose TPU Mosaic serialization pass to JAXlib
The JAXlib build now includes the TPU Mosaic serialization pass by forwarding the header from the XLA Mosaic dialect. This change allows JAXlib to access and utilize the serialization logic defined in the XLA codebase, ensuring that the latest version of the MosaicSerdePass is available within the JAXlib environment.
jaxlib/mosaic/dialect/tpu/transforms · high confidence
Finalize deprecation of zero-dimensional inputs to jnp.nonzero
The \jnp.nonzero\ function now raises an error when passed zero-dimensional (scalar) inputs, finalizing the deprecation of this usage. This change aligns the function's behavior with NumPy's handling of scalar inputs and ensures consistent error reporting for invalid argument shapes.
_jax/\src/numpy · high confidence
Internal restructuring of JAX sparse linear algebra solvers
The implementation of JAX's sparse linear algebra solvers (including CG and BiCGSTAB) has been moved from the public API into the internal module \jax.\_src.scipy.sparse.linalg\. This change consolidates the solver logic, including helper functions for matrix-vector products and tolerance handling, into a private source location, separating the internal implementation details from the public interface.
_jax/\src/scipy/sparse · high confidence
JAX 0.11.2 release
This update releases JAX version 0.11.2 and synchronizes the jaxlib version to 0.1.77. The release includes a comprehensive set of changes such as the deprecation of several APIs (including \jax.experimental.pjit.NamedSharding\, \jax.interpreters.pxla\ symbols, and various \jax.core\ functions), the removal of deprecated modules like \jax.abstract\_arrays\ and \jax.experimental.maps\, and the introduction of new features like \jax.scipy.stats.gaussian\_kde\ and \jax.lax.split\. It also addresses build and dependency updates, including upgrades to XLA, Bazel, and various Python tooling, alongside numerous bug fixes and performance improvements.
(repo-wide) · high confidence
JAX version bump to 0.10.0
The JAX library has been updated to version 0.10.0. This release includes the finalization of deprecations for several \jax.core\ APIs, the removal of deprecated symbols from \jax.dlpack\, \jax.errors\, \jax.lib.xla\_bridge\, \jax.lib.xla\_client\, and \jax.lib.xla\_extension\, and the removal of the \jax.\_src\ module from the public JAX namespace.
jax · high confidence
Lazy loading of MLIR dialects and introduction of StableHLO alias
The \jax.\_src.lib.mlir.dialects\ module now uses lazy loading for most MLIR dialects (such as \arith\, \builtin\, \cf\, \gpu\, \mhlo\, etc.), importing them from \jaxlib.mlir.dialects\ only when accessed, which improves startup performance. Additionally, the \stablehlo\ dialect is explicitly imported and aliased as \hlo\ to abstract the transition from MHLO to StableHLO, while \mpmd\ and \sdy\ dialects are loaded eagerly.
_jax/\src/lib/mlir/dialects · high confidence
MLIR context management and threading model changes
JAX now uses one MLIR context per Python thread instead of one per computation, and disables threading within MLIR contexts. This change improves thread safety and performance by avoiding the overhead of creating and destroying MLIR contexts for every computation and preventing race conditions when dialects are registered concurrently.
_jax/\src/interpreters · high confidence
Migration of JAX SciPy implementation to private source namespace
The JAX SciPy module has been moved from the public \jax.scipy\ namespace into the private internal package \jax.\_src.scipy\. This change consolidates the implementation files (such as \fft.py\, \linalg.py\, \signal.py\, \special.py\, and \ndimage.py\) under the private source tree, marking the end of the public \jax.scipy\ API and preparing for its removal in a future version. Users relying on the public \jax.scipy\ interface should be aware that this location is now considered internal implementation detail.
_jax/\src/scipy · high confidence
Mosaic GPU layout inference and lowering overhaul
The Mosaic GPU subsystem in \jax/experimental/mosaic/gpu\ has been significantly refactored to support a new equation-driven layout inference system and expanded lowering capabilities for Blackwell and Hopper architectures. Users benefit from improved automatic layout inference for shared memory (SMEM) and tensor memory (TMEM), enabling more efficient data transfers and reduced bank conflicts in matrix multiply-accumulate (MMA) operations. The update introduces support for warpgroup semantics, allowing for more flexible kernel scheduling and collective operations across clusters. Additionally, the profiler has been enhanced with CUPTI V2 multi-subscriber support and warp-level profiling, providing more accurate performance metrics. Bug fixes address issues with type conversions, barrier handling, and async copy operations, ensuring stability with newer CUDA toolkits and jaxlib versions.
jax/experimental/mosaic/gpu · high confidence
Mosaic TPU lowering is now a first-party JAX module
The Mosaic TPU lowering implementation has been moved from \jax.experimental.mosaic\ into the stable \jax.\_src.pallas.mosaic\ package. This change makes the Mosaic lowering rules a core part of JAX's Pallas infrastructure, removing the experimental dependency and ensuring they are built and tested alongside the main Pallas codebase.
_jax/\src/pallas/mosaic · high confidence
New API for defining custom PRNG implementations
The \jax.\_src.extend\ module now provides \jax.extend.random.define\_prng\_impl\, a function that allows users to register custom pseudo-random number generator (PRNG) implementations by specifying key shape, seed, split, random\_bits, and fold\_in behaviors. This replaces the previous \extend.random.PRNGImpl\ class-based approach, offering a more streamlined interface for extending JAX's random number generation capabilities.
_jax/\src/extend · high confidence
New Bazel-based build system and CLI for JAX
The build process has been restructured to use a new Bazel-based system with a dedicated CLI (\build/build.py\) and \build/BUILD.bazel\ targets. This change introduces hermetic Python support, allowing builds to specify a Python version explicitly, and centralizes dependency management through \requirements.in\ and \test-requirements.txt\ processed by \rules\_python\. The new system also adds support for building various wheel artifacts (including CUDA, ROCm, and OneAPI plugins) and editable installs, while providing utilities like \parallel\_accelerator\_execute.sh\ for test distribution and \update\_requirements\_uv.py\ for lockfile generation.
build · high confidence
New CUDA plugin extension and version introspection module
The \jaxlib/cuda\ directory now contains the core CUDA plugin extension (\cuda\_plugin\_extension\) and a new version introspection module (\\_versions\). The plugin extension exposes FFI types and handlers (including GPU callback and DCE sink handlers) and provides a \get\_device\_ordinal\ utility, while the version module exposes runtime and build-time version checks for CUDA libraries (cuBLAS, cuDNN, etc.) and device capabilities (compute capability, multicast support). This change also introduces a compile-time requirement for CUDA 11.8 or newer and switches the Python binding implementation from pybind11 to nanobind.
jaxlib/cuda · high confidence
New modular ROCm plugin package structure
The ROCm GPU plugin is now distributed as a dedicated, version-specific package (e.g., \jax\_rocm7\_plugin\) located in \jax\_plugins/rocm\. This change introduces a new build and initialization flow where the plugin detects the installed ROCm major version to load the correct extension, registers custom FFI types and call handlers for the 'ROCM' platform, and includes a warning for multi-GPU setups if shared memory is insufficient. Users benefit from clearer version isolation and explicit support for Python 3.12–3.14 in the plugin metadata.
_jax\plugins/rocm · high confidence
New modular random submodule with additional PRNG implementations
The random module has been reorganized into a new \jax.\_src.random\ submodule, restructuring the internal source layout. This change introduces new PRNG implementations—Philox 2x32, Philox 4x32, and Threefry 4x32—alongside the existing Threefry 2x32 and RngBitGenerator (RBG) variants, providing users with more options for random number generation. The \jax.\_src.random.core\ module now centralizes core random functions and key operations, while \jax.\_src.random.prng\ defines the PRNG implementation interface and key array classes.
_jax/\src/random · high confidence
Pallas API reorganization and deprecation cleanup
The \jax.experimental.pallas\ module has been reorganized to expose a unified public API while consolidating backend-specific functionality into dedicated submodules (\pallas.mosaic\_gpu\, \pallas.tpu\, \pallas.triton\, \pallas.tpu\_sc\, and \pallas.fuser\). This change introduces deprecation warnings for several legacy symbols, including \pl.core\_map\ (use \pl.kernel\), \pl.dot\ (removed in favor of \jax.numpy.dot\ or \einsum\), and \pl.reciprocal\ (moved to \pltpu\). The \pallas\_call\ primitive is now the standard entry point for kernel execution, and backend-specific compiler parameters and memory spaces are now accessed via their respective submodules.
jax/experimental/pallas · high confidence
Pallas ops module now registers its directory as user code for source info
The \jax/experimental/pallas/ops/\_\init\\_.py\ module now explicitly registers its own directory as a user-code location for JAX's source information tracking. This ensures that stack frames originating from Pallas operations are not filtered out during debugging or error reporting, allowing users to see full tracebacks and source context for Pallas-specific code.
jax/experimental/pallas/ops · high confidence
Public API reorganization and expansion in jax.lax
The \jax.lax\ module has been restructured to expose a curated public API while moving internal implementation details to private source paths. This change introduces numerous new primitives and functions—including \broadcast\_like\, \polynomial\, \integer\_pow\, \exp2\, \log2\, \mulhi\, \one\_minus\_square\, \reduce\_precision\, \ragged\_dot\, and \optimization\_barrier\—along with a new \jax.lax.linalg\ submodule for linear algebra operations (e.g., \cholesky\, \eigh\, \svd\, \tridiagonal\solve\). Simultaneously, many previously public internal helpers (prefixed with \\\, such as \\_delta\, \\_complex\, \\_reduce\_sum\) and deprecated symbols (e.g., \zeros\_like\_array\, \prod\, \infeed\, \outfeed\) have been removed from the public namespace to reduce clutter and enforce cleaner usage patterns.
jax/lax · high confidence
Python 3.15 support and improved Bazel build reliability
The build system now supports Python 3.15, including the free-threaded variant, by adding checksums for the CPython 3.15.0rc1 standalone builds and updating the \rules\_python\ configuration. To ensure reliable package resolution, patches were applied to \rules\_python\ to fix pure-Python wheel matching on Bazel modules (bzlmod) by including the 'none' ABI tag, and to support local wheel overrides for development workflows. Additionally, Windows build stability is improved by patching the Python site initialization to correctly handle long file paths via extended-length prefixes, and by bumping \rules\_python\ to version 2.2.0.
_third\party/py · high confidence
ROCm wheel builds migrate to manylinux and integrate Clang
The ROCm wheel build process now uses a manylinux\_2\_28 base image and installs LLVM 18 with Clang support, replacing previous build environments. This change includes a new configuration file (clang.cfg) that directs Clang to use the system GCC toolchain for standard libraries, ensuring compatibility during compilation. Additionally, the build environment now explicitly installs numactl-devel to resolve library linking issues and supports a broader range of GPU architectures, including new gfx12xx targets.
_build/rocm/build\wheels · high confidence
Refactored build CLI with subcommands and hermetic Bazel 9.2.0 support
The build tooling has been restructured to use a subcommand-based CLI architecture, moving utility functions into dedicated modules (command.py, utils.py) and introducing a CommandBuilder/SubprocessExecutor for managing subprocess execution with optional detailed timestamped logging. This change also upgrades the default Bazel version to 9.2.0, including hermetic toolchain support for Linux aarch64 builds, and updates the build scripts to reflect these structural and dependency changes.
build/tools · high confidence
Refactored image resizing implementation into a dedicated source module
The internal implementation of image resizing functions has been moved from the public \jax.image\ namespace into a new private module at \jax.\_src.image\. This change includes the introduction of new interpolation methods, specifically AREA resizing and PYTORCH\_CUBIC, and adds documentation clarifying the output data type of \jax.image.resize\. The refactoring also ensures consistent handling of negative dimensions in scaling operations and prevents NaN values in the Lanczos kernel and scaling logic.
_jax/\src/image · high confidence
Restructure CUDA plugin into a dedicated package with explicit library loading
The CUDA plugin code has been moved from jaxlib into a new jax\_plugins/cuda package, introducing a standalone build (BUILD.bazel) and Python packaging (setup.py, pyproject.toml) for the plugin. This change restructures how the plugin is distributed and loaded: the plugin now explicitly loads NVIDIA libraries (cuBLAS, cuDNN, etc.) up front via ctypes to ensure they are found, and the Python entry point is registered via jax\_plugins entry points rather than relying on the previous jaxlib internal extension structure.
_jax\plugins/cuda · high confidence
Restructured internal library module with new Mosaic GPU and Triton bindings
The \jax.\_src.lib\ package has been reorganized into a dedicated source location (\jax/\_src/lib\) with a new Bazel build file, consolidating imports from \jaxlib\ and exposing updated bindings. This change introduces Python bindings for the Mosaic GPU MLIR dialect (via \mosaic\_gpu.py\) and the Triton MLIR dialect (via \triton.py\), while also updating the internal profiler interface (\\_profiler.pyi\) and refining CUDA path detection logic.
_jax/\src/lib · high confidence
Updated gRPC integration with Bazel plugin flag support
The gRPC third-party dependency has been updated to version 1.81.0. This change introduces support for passing custom flags to the gRPC code-generation plugin via a new \plugin\_flags\ argument in the \cc\_grpc\_library\ rule. Additionally, the build configuration has been cleaned up by removing temporary include paths for \upb-gen\ and \upbdefs-gen\ that were previously marked as workarounds for code-generation issues.
_third\party/grpc · high confidence
Fixes
Fix crash returning a Token from a jit computation on GPU
Resolves a crash that occurred when a JIT-compiled function returned a Token on GPU devices. This change ensures that Token values are handled correctly during JIT execution on GPU, preventing runtime errors in computations that rely on ordered effects or token passing.
jax/interpreters · high confidence
Fixes for JAX Bzlmod builds
This change resolves build issues with JAX when using Bazel's Bzlmod system. It introduces a new BUILD.bazel file for the third-party protobuf dependency, applies a patch to work around compiler and static analysis complaints regarding incomplete types in message\_lite.h, and adds a public alias for the arena target to ensure proper visibility.
_third\party/protobuf · high confidence
Test coverage
Added Colab test notebooks for CPU, GPU, and TPU; Added Flax model examples for jax2tf testing; Added Pallas export backward compatibility test data; Added TSAN test configuration and wrapper for RBE environments; Added backward compatibility test data for tf.call\_tf\_function; Added multiprocess test for JAX2TF multihost export; Added tests for FFI example modules; Added tests for JAX scipy line\_search compatibility; Added tests for Pallas Triton ragged\_dot lowering on GPU; Expanded test coverage for JAX core and experimental features; Expanded test coverage for jax2tf conversion and shape polymorphism; Extensive Pallas test suite expansion and maintenance; Initial test suite for Mosaic GPU dialect and layout inference; New FileCheck tests for JAX MLIR lowering; New ULP precision test suite for JAX numerics; New export backward-compatibility test data for device placement and LAPACK solvers; Open-source multiprocess test suite and runner; Removal of legacy JAX test suite files.
Dependencies
Regenerated Python 3.12, 3.13, and 3.14 dependency lock files
The hermetic Python dependency lock files for Python 3.12, 3.13, and 3.14 (including the free-threaded 3.14 variant) have been regenerated using \uv\. This update refreshes the pinned versions and cryptographic hashes for all transitive dependencies, ensuring that builds for these Python versions use the latest compatible package sets defined in the source requirements.
(dependencies) · high confidence
Upgraded Abseil to LTS 20260526.0 with build fixes
The Abseil library has been upgraded to the LTS 20260526.0 release, which includes patches to improve Bazel build compatibility and fix compiler header usage. Specifically, the btree iterator type definition has been simplified to remove conditional logic, and the raw\_hash\_set implementation now includes immintrin.h for x86/x86\_64 platforms instead of bmi2intrin.h to resolve build issues.
_third\party/absl · high confidence
Housekeeping
Regenerated Python type stubs for jaxlib.\_jax
The Python type stubs (.pyi files) in the jaxlib.\_jax package have been regenerated to align with the current C++ bindings. This update refreshes the type signatures for core components including the JaxRuntimeError exception, CompileOptions, the pytree registry and definitions, and the guard libraries, ensuring that static type checkers and IDEs receive accurate information about the available APIs and their expected argument and return types.
_jaxlib/\jax · high confidence
Written by watchdog.canine.dev from the codebase's own history, inside the signed delivery this page is composed from.
How this codebase got here
Score
- CAI 51 → 68 (+17.0)
- Rubric changed (rubric-2026.08.15 → rubric-2026.09.15) — scores are not directly comparable.
Lenses
- Code Health 100 → 77 (-22.9)
- Architecture 98 (new)
- Maturity 73 → 69 (-4.3)
- Readiness 32 → 72 (+39.9)
- Security 58 → 63 (+4.5)
- Accessibility 74 (new)
Resolved (23)
- Coverage not measured — test suite did not build
- Dimension evaluation failed
- High IaC: DS-0029 (build/rocm/Dockerfile.ms)
- High IaC: DS-0029 (build/rocm/Dockerfile.ms)
- High IaC: DS-0029 (build/rocm/Dockerfile.ms)
- High IaC: DS-0029 (build/rocm/docker/Dockerfile.jax-ubu22)
- High IaC: DS-0029 (build/rocm/docker/Dockerfile.jax-ubu22)
- High IaC: DS-0029 (build/rocm/docker/Dockerfile.jax-ubu24)
- High IaC: DS-0029 (build/rocm/docker/Dockerfile.jax-ubu24)
- High IaC: KSV-0118 (ci/k8s/indexed-job.yaml)
- High IaC: KSV-0118 (ci/k8s/indexed-job.yaml)
- Low IaC: DS-0026 (build/rocm/Dockerfile.ms)
- Low IaC: DS-0026 (build/rocm/docker/Dockerfile.jax-ubu22)
- Low IaC: DS-0026 (build/rocm/docker/Dockerfile.jax-ubu24)
- Low IaC: KSV-0004 (ci/k8s/indexed-job.yaml)
- Low IaC: KSV-0020 (ci/k8s/indexed-job.yaml)
- Low IaC: KSV-0021 (ci/k8s/indexed-job.yaml)
- Low IaC: KSV-0030 (ci/k8s/indexed-job.yaml)
- Low IaC: KSV-0106 (ci/k8s/indexed-job.yaml)
- No exposed public API
- …and 3 more
New (2926)
- ArrayImpl.__dlpack_device__ (cognitive 21) (jax/_src/array.py)
- ArrayImpl._value (cognitive 25) (jax/_src/array.py)
- Barrier._wait (cognitive 40) (jax/_src/pallas/mosaic_gpu/interpret/shared_memory.py)
- Barrier._wait (cyclomatic 26) (jax/_src/pallas/mosaic_gpu/interpret/shared_memory.py)
- Barrier.arrive (cognitive 18) (jax/_src/pallas/mosaic_gpu/interpret/shared_memory.py)
- BarrierRef.arrive (cognitive 17) (jax/experimental/mosaic/gpu/utils.py)
- BlockSpec.to_block_mapping (cognitive 33) (jax/_src/pallas/core.py)
- BlockSpec.to_block_mapping (cyclomatic 24) (jax/_src/pallas/core.py)
- BufferedRef.create (cognitive 36) (jax/_src/pallas/mosaic/pipeline.py)
- BufferedRef.create (cyclomatic 21) (jax/_src/pallas/mosaic/pipeline.py)
- CI installs an unverified third-party binary (.github/workflows/build_rocm_artifacts.yml)
- CI installs an unverified third-party binary (.github/workflows/tsan.yaml)
- Change coupling clique: serialization.py, export_with_memory_space.py, export_with_specified_sharding.py, export_with_unspecified_sharding.py (jax/_src/export/serialization.py)
- Change coupling: abstract_arrays.py ↔ mlir.py (jax/_src/abstract_arrays.py)
- Change coupling: compute_on.py ↔ fused.py (jax/_src/compute_on.py)
- Change coupling: conditionals.py ↔ fuser_utils.py (jax/_src/lax/control_flow/conditionals.py)
- Change coupling: conditionals.py ↔ solves.py (jax/_src/lax/control_flow/conditionals.py)
- Change coupling: jet.py ↔ transform.py (jax/experimental/jet.py)
- Change coupling: lowering.py ↔ torch.py (jax/_src/pallas/mosaic_gpu/lowering.py)
- Change coupling: pjit.py ↔ fused.py (jax/_src/pjit.py)
- …and 2906 more
Changes since last survey
- 300 commits — 269 feature/other, 31 fixes
By area
- jax/_src — 109 commits
- (repo) — 64 commits
- tests/pallas — 31 commits
- (root) — 29 commits
- jax/experimental — 16 commits
- jaxlib/mosaic — 7 commits
- tests/numerics — 5 commits
- docs/accuracy.md — 3 commits
- tests/memories_test.py — 3 commits
- docs/101 — 2 commits
- jaxlib/BUILD — 2 commits
- jaxlib/gpu — 2 commits
- jaxlib/xla_client.py — 2 commits
- tests/ann_test.py — 2 commits
- tests/layout_test.py — 2 commits
- tests/mosaic — 2 commits
- .github/workflows — 1 commit
- benchmarks/tracing_benchmark.py — 1 commit
- ci/envs — 1 commit
- docs/201 — 1 commit
Notable commits
- fix: Add large buffer regression test for psum_scatter.
- fix: Bump rules_ml_toolchain to fix CUDA build under Bazel 9.2.0.
- fix: Fix Windows frexp exponent normalization, TPU 7x chip count, and guard recent SparseCore tests.
- fix: Fix testSincValuesAndDerivativesLargeMagnitude with older NumPy versions.
- fix: Fix a CI test failure on Windows.
- fix: Fix a data race by removing redundant flags.writeable assignment in Array._value.
- fix: Fix build breakage.
- fix: Fix eager partial-manual shard_map with check_vma=False
- fix: Fix erf accuracy test bounds on GPU by disabling input_ftz.
- fix: Fix expm1 gradient cancellation at large-negative inputs
- fix: Fix float16 tolerance in testReducerInitial
- fix: Fix jnp.sinc gradient and higher-order derivative accuracy near zero
- fix: Fix out-of-bounds slice check in NDIndexer when start is 0.
- fix: Fix partition, argpartition and top_k on signed integer minimums
- fix: Fix the test failures caused by removing the fully-replicated check.
- fix: Fix transpose rule for lax.select_n with out-of-range indices or single case
- fix: Fixed test_emit_pipeline_small on TPU v3
- fix: Merge pull request #40700 from ROCm:fix/reducer-initial-f16-tolerance
- fix: Merge pull request #40820 from Phoebus-Liu:fix-39794
- fix: Merge pull request #40961 from AshishKumar4:fix-eager-shardmap-check-vma
- …and 280 more
Written by watchdog.canine.dev from the codebase's own history, inside the signed delivery this page is composed from.
Survey your own repository
jax-ml/jax was measured the same way every project in this corpus was: the same rubric, at a pinned commit, with the result published in full. Point a surveyor at a repository you know and see whether you agree with it.
About this page
- The score is its most recent published measurement, taken on 26 September 2026 at a pinned commit. It is not a live figure and does not change until the project is measured again.
- Measured at commit 886d2370c1c959d210f522e352e3ddc6bcff7d6c — the exact code this score is about.
- Scored under rubric-2026.09.15 — the same rubric and the same method as every other entry in this index.
- Measured by watchdog.canine.dev using codehealth-analyzer preprod-d0929f7ac71f.