jax-ml/jax

Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more

35,442 stars Python 10 components

13 hidden assumptions · 6-stage pipeline · 10 components

Like any codebase, this ml training makes assumptions it never checks — most are routine. The ones worth your attention are below.

Compiles NumPy programs with automatic differentiation to accelerated hardware via XLA

User Python functions are traced into Jaxpr intermediate representation, transformed by interpreters (differentiation, vectorization, compilation), then lowered to MLIR and compiled to XLA for execution on accelerated hardware. Each transformation creates new Jaxprs that preserve computation semantics while adding capabilities like gradients or batching.

Under the hood, the system uses 2 feedback loops, 3 data pools, 3 control points to manage its runtime behavior.

A 10-component ml training. 1234 files analyzed. Data flows through 6 distinct pipeline stages.

Hidden Assumptions

Most of what this code assumes is routine. These 3 are the ones most likely to cause trouble here. The rest are minor; they're under "Show everything".

Worth your attention first

Silent failures or crashes deep in XLA when jaxlib C++ extensions have incompatible ABIs, making debugging extremely difficult

Worth your attention first

GPU kernels may fail at runtime with cryptic CUDA errors if the wrong plugin version is loaded for the actual hardware/driver combination

Worth your attention first

ctypes.cdll.LoadLibrary succeeds but pycapsule creation fails with segfaults or undefined symbols at FFI call time

Show everything (10 more)
Shape

Assumes input arrays a and b have identical shapes and dtype float32, but only validates dtype and shape equality without checking memory layout, strides, or contiguity that the underlying CUDA kernel expects

If this fails: CUDA kernel may read incorrect memory locations or crash if arrays have non-contiguous strides or unexpected memory layout

examples/ffi/src/jax_ffi_example/cuda_examples.py:foo_fwd
Resource

Assumes cudaMemcpyAsync will complete before the CUDA stream is used by subsequent operations, but doesn't add explicit synchronization or check for CUDA errors from the memcpy operation

If this fails: Race conditions where downstream GPU operations read uninitialized data if memcpy hasn't completed, leading to wrong results

examples/ffi/src/jax_ffi_example/gpu_examples.cc:StateExecute
Domain

Assumes the input array attribute contains int32_t elements that can be safely summed into an int64_t without overflow, but never checks the array size or validates that sum won't exceed int64_t max value

If this fails: Integer overflow in the sum calculation leads to wrong results when processing large arrays or arrays with large values

examples/ffi/src/jax_ffi_example/cpu_examples.cc:ArrayAttrImpl
Ordering

Assumes input and output arrays have the same memory layout and that it's safe to read from x and write to y in the same sequential order, but doesn't validate pointer alignment or check for memory overlap

If this fails: Memory corruption or undefined behavior if input and output arrays overlap in memory or have different alignment requirements

examples/ffi/src/jax_ffi_example/rms_norm.cc:ComputeRmsNorm
Scale

Assumes the size parameter fits in int64_t and that the variance calculation (sm / size) won't cause numerical instability, but doesn't validate size > 0 or handle the case where all elements are zero

If this fails: Division by zero or numerical overflow when size is 0 or when sum of squares is extremely large, leading to NaN or Inf results

examples/ffi/src/jax_ffi_example/rms_norm.cc:ComputeRmsNorm
Contract

Assumes the num parameter is a positive integer that can be used to create a valid numpy array range, but never validates num >= 0 or checks that np.arange(num) won't exhaust memory

If this fails: Memory exhaustion or negative array indices when num is very large or negative, causing crashes or wrong results

examples/ffi/src/jax_ffi_example/cpu_examples.py:array_attr
Contract

Assumes the input x is a JAX array with a numeric dtype compatible with the C++ implementation, but never validates that x.dtype maps to a supported C++ type (T in the template)

If this fails: Type mismatch between Python array dtype and C++ template instantiation leads to incorrect memory interpretation and wrong results

examples/ffi/src/jax_ffi_example/rms_norm.py:rms_norm
Environment

Assumes jaxlib.mlir.dialects module exists and contains all expected dialect submodules (arith, mhlo, etc.), but uses lazy loading without validating module completeness

If this fails: AttributeError at runtime when trying to access missing dialect modules, breaking MLIR lowering for certain operations

jax/_src/lib/mlir/dialects/__init__.py:lazy_loading
Temporal

Assumes the State object lifetime extends beyond the StateExecute call and that the memory pointed to by state->value remains valid during the asynchronous CUDA memcpy

If this fails: Use-after-free or reading invalid memory if the State object is destroyed before the CUDA stream completes the memcpy operation

examples/ffi/src/jax_ffi_example/gpu_examples.cc:State
Domain

Assumes eps=1e-5 is an appropriate epsilon value for numerical stability across all input scales and dtypes, but doesn't adjust epsilon based on the working precision (float32 vs float64)

If this fails: Numerical instability for very small values with float32 or insufficient precision for float64 computations

examples/ffi/src/jax_ffi_example/rms_norm.py:eps_parameter

Open the standalone hidden-assumptions report for jax →

How Data Flows Through the System

User Python functions are traced into Jaxpr intermediate representation, transformed by interpreters (differentiation, vectorization, compilation), then lowered to MLIR and compiled to XLA for execution on accelerated hardware. Each transformation creates new Jaxprs that preserve computation semantics while adding capabilities like gradients or batching.

  1. Function Tracing — JaxprTrace intercepts Python function execution by replacing concrete values with Tracer objects, recording each operation as JaxprEqn entries to build a Jaxpr computation graph [Python function → Jaxpr]
  2. Abstract Evaluation — Each primitive operation computes abstract values (ShapedArray) to track shapes and types through the computation without executing on concrete data [JaxprEqn → ShapedArray]
  3. Interpreter Transform — Specialized interpreters like VJPTrace, BatchTrace, or JVPTrace transform the input Jaxpr by adding new equations for gradients, vectorization, or other transformations [Jaxpr → Jaxpr]
  4. MLIR Lowering — lower_jaxpr_to_module converts Jaxpr equations to MLIR operations using MlirLoweringContext, mapping JAX primitives to platform-specific MLIR dialects [Jaxpr → MLIR module]
  5. XLA Compilation — MLIR module is compiled through XLA to generate optimized machine code for the target accelerator (CPU, GPU, TPU), producing an XlaExecutable [MLIR module → XlaExecutable]
  6. Hardware Execution — XlaExecutable.call() executes compiled code on the target device with concrete ArrayImpl data, managing memory transfers and device synchronization [XlaExecutable → ArrayImpl]

Data Models

The data structures that flow between stages — the contracts that hold the system together.

Jaxpr jax/_src/core.py
NamedTuple with jaxpr: Jaxpr, literals: list[Any], consts: list[Array] — captures a computation graph with input variables (invars), output variables (outvars), and equations (eqns)
Created during function tracing by converting Python operations into primitive equations, transformed by interpreters (differentiation, batching), then lowered to MLIR
JaxprEqn jax/_src/core.py
NamedTuple with primitive: Primitive, invars: list[Var], outvars: list[Var], params: dict[str, Any] — represents a single operation in the computation graph
Generated when tracing Python operations, processed sequentially by interpreters to create new equations with transformed variables
ShapedArray jax/_src/abstract_arrays.py
AbstractValue with shape: tuple[int, ...], dtype: np.dtype, weak_type: bool — abstract representation of array shape and type without data
Computed during tracing to track array shapes through transformations, used for type checking and memory allocation planning
ArrayImpl jax/_src/array.py
JAX array implementation with aval: ShapedArray, sharding: Sharding, _arrays: list[Array] — concrete array data distributed across devices
Created from NumPy arrays or computation results, manages data placement across devices, destroyed when computation completes
XlaExecutable jax/_src/interpreters/pxla.py
Compiled executable with xla_executable: Any, input_avals: Sequence[ShapedArray], output_avals: Sequence[ShapedArray], input_shardings: Sequence[JSharding] — ready-to-run compiled function
Created by compiling Jaxpr through MLIR to XLA, cached for repeated execution, invoked with concrete array data

System Behavior

How the system operates at runtime — where data accumulates, what loops, what waits, and what controls what.

Data Pools

Compilation Cache (cache)
Caches XlaExecutable objects keyed by function signature and static arguments to avoid recompilation on repeated jit calls
Device Memory (buffer)
GPU/TPU memory pools managed by XLA for array storage and computation, handles allocation and deallocation of device buffers
Primitive Registry (registry)
Global registry mapping primitive names to Primitive objects with their transformation rules for each interpreter

Feedback Loops

Delays

Control Points

Technology Stack

XLA (compute)
Compiles MLIR to optimized machine code for accelerators, handles memory management and device scheduling
MLIR (framework)
Multi-level intermediate representation for compiler transformations, bridges JAX primitives to hardware-specific operations
NumPy (library)
Provides array API compatibility and reference semantics for JAX array operations
jaxlib (runtime)
C++ extension providing XLA bindings, CUDA/TPU support, and low-level array operations

Key Components

Explore the interactive analysis

See the full architecture map, data flow, and code patterns visualization.

Analyze on CodeSea

Related Ml Training Repositories

Frequently Asked Questions

What is jax used for?

Compiles NumPy programs with automatic differentiation to accelerated hardware via XLA jax-ml/jax is a 10-component ml training written in Python. Data flows through 6 distinct pipeline stages. The codebase contains 1234 files.

How is jax architected?

jax is organized into 5 architecture layers: User API, Core Tracing, Interpreters, MLIR Lowering, and 1 more. Data flows through 6 distinct pipeline stages. This layered structure keeps concerns separated and modules independent.

How does data flow through jax?

Data moves through 6 stages: Function Tracing → Abstract Evaluation → Interpreter Transform → MLIR Lowering → XLA Compilation → .... User Python functions are traced into Jaxpr intermediate representation, transformed by interpreters (differentiation, vectorization, compilation), then lowered to MLIR and compiled to XLA for execution on accelerated hardware. Each transformation creates new Jaxprs that preserve computation semantics while adding capabilities like gradients or batching. This pipeline design reflects a complex multi-stage processing system.

What technologies does jax use?

The core stack includes XLA (Compiles MLIR to optimized machine code for accelerators, handles memory management and device scheduling), MLIR (Multi-level intermediate representation for compiler transformations, bridges JAX primitives to hardware-specific operations), NumPy (Provides array API compatibility and reference semantics for JAX array operations), jaxlib (C++ extension providing XLA bindings, CUDA/TPU support, and low-level array operations). A lean dependency footprint.

What system dynamics does jax have?

jax exhibits 3 data pools (Compilation Cache, Device Memory), 2 feedback loops, 3 control points, 2 delays. The feedback loops handle training-loop and cache-invalidation. These runtime behaviors shape how the system responds to load, failures, and configuration changes.

What design patterns does jax use?

3 design patterns detected: Transformation Composition, Trace-Transform-Lower, Abstract Interpretation.

Analyzed on April 20, 2026 by CodeSea. Written by .