XLA (Accelerated Linear Algebra) is the compiler underneath JAX, and also behind PyTorch/XLA and other frameworks that target TPUs. It takes a whole numerical program with static shapes, optimises it as one graph, and produces a single executable for a GPU, TPU or CPU. If you train or serve with JAX, every jax.jit call goes through it. When a step is fast, it is usually because XLA fused the work, chose good memory layouts, or split it well across devices. When the first step takes four minutes, or every tenth step stalls, it is usually because XLA is compiling again.

This article follows a function from Python to device code and explains what each stage does and how it affects training. It then shows how to look at XLA's output yourself, works through an MLP block and what fusion does to its memory traffic, and covers the recompilation and memory failures that cause most problems in practice. The PyTorch compiler path is covered separately in the torch.compile deep dive.

The pipeline, stage by stage

Tracing. When you call a jitted function, JAX runs your Python once with tracer objects instead of real arrays. A tracer carries only a shape and a dtype. Every jnp operation records a primitive in a jaxpr, a small typed program with no Python left in it. Python control flow runs while tracing, so a loop over a Python range gets unrolled, and an if on a traced value fails with a tracer conversion error because the value is not known yet. Control flow that depends on data has to use lax.cond, lax.while_loop or lax.scan, which become real control-flow operations in the compiled program.

Lowering. The jaxpr is lowered to StableHLO, an MLIR dialect with a versioned, portable set of operations (dot_general, convolution, reduce, broadcast, while, and so on). StableHLO is the boundary between frameworks and compilers in the OpenXLA project, which is why the same model can be exported and compiled elsewhere.

HLO optimisation. XLA converts the program to its internal HLO and runs a long series of passes. Algebraic simplification and common-subexpression elimination remove redundant work. Sharding propagation and the SPMD partitioner (Shardy, which has replaced GSPMD as JAX's propagation system) turn a program written for one large array into a per-device program with explicit collectives. Fusion merges operations into kernels. Layout assignment chooses how each tensor is laid out in memory. Scheduling and buffer assignment decide the order of operations and which tensors can share memory.

Backend code generation. On GPUs, fused elementwise and reduction kernels are emitted through LLVM to PTX. Matrix multiplies and convolutions either go to vendor libraries or, for some fusions, to Triton-generated kernels. Autotuning picks among the candidates by timing them on the device during compilation, which is one reason GPU compiles take a while. TPU and CPU have their own backends. The result is wrapped as a PJRT executable. PJRT is the plug-in runtime interface that lets device vendors connect a backend without changing the framework.

From a Python function to device codePython + jax.jityour functionjaxprtraced on abstract shapesStableHLOportable MLIR dialectHLO moduleXLA's own IRtracelowerTarget-independent and target-specific HLO passessimplify and CSE | sharding propagation and SPMD partitioning | fusionlayout assignment | scheduling | buffer assignment (liveness, aliasing)GPU backendLLVM to PTX, library and Triton GEMMsTPU / CPU backendstarget code generationPJRT executablecached by shapes, dtypes, sharding
jax.jit traces Python to a jaxpr, lowers it to StableHLO, and XLA optimises the HLO module as a whole before backend code generation. The executable is cached under a key built from shapes, dtypes, static arguments and sharding.

Fusion, layout and buffers

On modern accelerators, most operations in a model are limited by memory bandwidth, not compute. An elementwise add does one floating-point operation per element but reads two values and writes one. A chain of such operations, run one kernel at a time, sends every intermediate tensor out to HBM and reads it back again. Fusion puts the chain into one kernel, so intermediates stay in registers and the tensor goes through HBM once. XLA has several kinds of fusion: loop fusion of elementwise operations, input fusion of a reduction with the elementwise work that feeds it, and output fusion of elementwise work into the epilogue of a producer such as a matrix multiply. The general ideas are covered in the kernel fusion article. What makes XLA different is that it sees the whole step at once, forward and backward pass together, so it can fuse across the boundaries a layer-by-layer framework would stop at.

Fusion has limits. XLA will not fuse an expensive producer into several consumers if that means recomputing it many times. It will not create a kernel too big to fit the hardware's registers or shared memory. Some operations, such as large matrix multiplies and sorts, often end up in their own kernels. Whether a given chain was fused is something to check in the optimised HLO, not something to assume.

Layout and buffers. Layout assignment picks the physical order of each tensor's dimensions (the {1,0} annotations in HLO text) to suit the consumer, and inserts transposes or copies when two consumers want different layouts. Buffer assignment uses liveness analysis to let tensors whose lifetimes do not overlap share memory. This is why XLA's peak memory is often well below the sum of all intermediate sizes, and why donate_argnums matters: donating an input lets XLA write the output (for example, updated parameters) into the input's buffer instead of holding both.

Looking inside: the AOT API, dumps and compile logs

JAX lets you stop at each stage and look. The ahead-of-time API splits jit into lower and compile steps, and each result can be printed or analysed:

import jax, jax.numpy as jnp

def block(x, w, b):
    h = jax.nn.gelu(x @ w + b)
    mu = h.mean(-1, keepdims=True)
    var = ((h - mu) ** 2).mean(-1, keepdims=True)
    return (h - mu) * jax.lax.rsqrt(var + 1e-5)

x = jnp.ones((8192, 4096), jnp.bfloat16)
w = jnp.ones((4096, 4096), jnp.bfloat16)
b = jnp.zeros((4096,), jnp.bfloat16)

print(jax.make_jaxpr(block)(x, w, b))       # the traced program
lowered = jax.jit(block).lower(x, w, b)
print(lowered.as_text())                     # StableHLO handed to XLA
compiled = lowered.compile()
print(compiled.as_text())                    # optimised HLO: look for fusion ops
print(compiled.cost_analysis())              # XLA's estimate of flops and bytes
print(compiled.memory_analysis())            # argument, output and temp sizes

The structure returned by cost_analysis() has changed between JAX versions, so read it interactively before writing code that depends on it. For a full picture, set XLA_FLAGS=--xla_dump_to=/tmp/xla_dump before the process starts. XLA then writes each module's HLO before and after optimisation to that directory, so you can compare them and see which operations were fused, what layouts were picked and where collectives were added. For compile behaviour, JAX documents two configuration options worth turning on in development:

jax.config.update("jax_log_compiles", True)         # log every trace/lower/compile
jax.config.update("jax_explain_cache_misses", True) # say WHY a cache missed
jax.config.update("jax_compilation_cache_dir", "/mnt/cache/jax")  # persistent cache

The persistent cache saves compiled executables to disk, so a restarted job or a second worker with the same program, shapes, JAX version and hardware skips compilation. Set it before the first compile. To time device work correctly, call block_until_ready() on the result. JAX dispatches asynchronously, so without it you measure only how long it took to queue the work.

Worked example: an MLP block's memory traffic

Take the block function above: an 8192 by 4096 activation in bfloat16, a 4096 by 4096 weight, then bias, GELU and a layer normalisation. Each 8192 by 4096 bfloat16 tensor is 64 MiB. The matrix multiply does about 2 x 8192 x 4096 x 4096, or about 275 GFLOP, which keeps a modern GPU's tensor cores busy for well under a millisecond. Everything after it is memory traffic.

Run eagerly, one kernel per operation, the tail is roughly: bias add (read 64, write 64), GELU, which is several elementwise operations in its tanh form (each about 128 MiB of traffic), mean (read 64), subtract, square, mean, then the final subtract, multiply and rsqrt. That adds up to more than a gigabyte of HBM traffic for a few hundred MiB of real data. At roughly 3 TB/s that is about 0.4 ms, which can be more than the matrix multiply itself.

In the optimised HLO you would typically see the bias and GELU fused with the producer or into a single loop fusion, and the layer norm turned into one or two fusions built around its reductions. The tail then reads the matrix-multiply output once or twice and writes the result once: about 128 to 192 MiB instead of more than a gigabyte. The exact fusion boundaries depend on the backend and XLA version, which is why you check the dump instead of predicting it. Inside a full training step, the backward pass of this block gets the same treatment, and buffer assignment decides which forward activations stay alive for the backward pass. That is where jax.checkpoint (rematerialisation) comes in, trading recomputation for memory.

Recompilation: the silent slowdown

An XLA executable is specialised to its input shapes, dtypes, static arguments, sharding and pytree structure. Change any of these and jit traces and compiles again. On a large model a compile can take minutes, so recompilation is the most common XLA performance problem, and it does not look like an error. It looks like a slow step.

  • Variable sequence lengths. Each new length compiles a new program. Pad to a small set of buckets (for example powers of two, or multiples of 128) and mask the padding.
  • Ragged last batch. The final batch of an epoch is smaller and triggers a compile. Drop it, or pad it and mask the loss.
  • Python values as static arguments. A learning rate or step counter passed through static_argnums compiles once per value. Pass it as an array instead.
  • Dtype drift. A float32 batch where bfloat16 was expected, or a changing weak type from a Python scalar, is a different signature.
  • Unrolled Python loops. A 48-layer model written as a Python loop produces a program 48 times as large, and compile time grows with it. Use lax.scan over stacked layer weights to compile the layer once.

With jax_explain_cache_misses enabled, each miss is logged with the argument that changed, which turns hours of guessing into reading a log line.

SPMD partitioning

For multi-device training you annotate shardings on inputs (and sometimes on intermediates), and the partitioner works out the rest. It propagates shardings through the graph, rewrites each operation into its per-device form, and inserts all-gather, reduce-scatter, all-reduce or all-to-all where data has to move. The result is one program that every device runs on its own shard. Two practical points follow. First, the collectives XLA inserts are visible in the optimised HLO, so when a step is slower than expected, look there for an unexpected all-gather of a large weight. Second, a missing or contradictory annotation is rarely an error; it usually just leads to a sharding you did not intend. Hardware-specific interconnect behaviour on TPU is covered in the TPU article.

Failure modes

  • Slow first step. Expected, but long compiles in production restarts are avoidable: turn on the persistent cache, or compile ahead of time with lower(...).compile() during startup before taking traffic.
  • Periodic stalls. Almost always recompilation. Turn on compile logging and cache-miss explanations, then bucket shapes.
  • Out of memory at compile or first run. Read memory_analysis() for the temp size. Add rematerialisation, donate parameter and optimiser buffers, or shard further.
  • Numerical differences from eager. Fusion and algebraic simplification can reorder floating-point operations, and default matmul precision on some hardware uses reduced-precision inputs. Compare with tolerances, and set precision explicitly where it matters.
  • Host round-trips. print, .item() or converting to NumPy inside the step forces synchronisation. Use jax.debug.print inside jitted code and pull metrics out every N steps.
  • Custom kernels. When XLA's fusion is not good enough for a critical operation (attention variants, for example), JAX's Pallas extension lets you write a kernel by hand. It is a specialist tool, comparable to writing a Triton kernel.

Trade-offs

ChoiceGainCost
Whole-program compile vs eagerCross-op fusion, global buffer reuseCompile time; static shapes required
Shape bucketingBounded number of compilesWasted compute on padding
lax.scan over layersCompile once per layer typeLess per-layer scheduling freedom
RematerialisationLower activation memoryExtra forward compute (often 20-30%)
GPU autotuningFaster kernelsLonger compiles, run-to-run variance in compile time
Persistent cacheNear-instant restartsInvalidated by JAX, XLA or driver upgrades

What to do next

  1. Turn on jax_log_compiles and jax_explain_cache_misses in development, and count compiles per run. The target is a small fixed number.
  2. Bucket sequence lengths and batch sizes, and keep step counters and hyperparameters as arrays, not static arguments.
  3. Replace Python layer loops with lax.scan over stacked weights.
  4. Dump one training step with --xla_dump_to and read the optimised HLO: check the fusions, layouts and collectives you expect.
  5. Check memory_analysis() before scaling up, and use donation and jax.checkpoint to fit.
  6. Configure jax_compilation_cache_dir on shared storage for multi-worker jobs and restarts.
  7. Benchmark with block_until_ready(), excluding the compile step.
Key takeaway: XLA compiles a whole numerical program with static shapes into one optimised executable. Its speed comes from fusion, layout choice, buffer reuse and automatic partitioning; its costs are compile time and its sensitivity to shape changes. Inspect what it produced instead of guessing, keep the set of shapes small, cache compiled programs, and the compiler will do most of the performance work for you.