All essays
BenchmarkCOMPARISONFEB 2026

JAX GPU Infrastructure: XLA Compilation, TPU vs GPU, and JIT Challenges

Technical deep dive into JAX GPU deployment: XLA compilation pipeline, TPU vs GPU tradeoffs, JIT compilation overhead, sharded device memory, pjit partitions, and production inference patterns on H100 and A100.

01

THE XLA COMPILATION PIPELINE ON GPU

JAX compiles Python functions into XLA (Accelerated Linear Algebra) HLO (high-level optimizer IR) which is then lowered to LLVM IR and finally to GPU PTX assembly. The pipeline operates in stages: `jax.jit(f)` traces the function by running it with example inputs (abstract shapes), captures the JAXPR computation graph, and passes it to XLA. XLA performs optimization passes: operation fusion (merging element-wise ops into single kernels), buffer assignment (reducing memory allocations via in-place updates), and layout assignment (choosing optimal tensor memory layouts for GPU). The compiled result is cached in `JAX_COMPILATION_CACHE` (default: `~/.cache/jax/`) keyed by the function's computational hash. For a Llama 70B layer, compilation takes 18-45 seconds on first invocation, depending on the complexity of the fusion optimizations applied.

The critical GPU-specific compilation behavior is the XLA fusion granularity setting via `XLA_FLAGS="--xla_gpu_autotune_level=4 --xla_gpu_memory_limit_slop_factor=100"`. The autotune level (1-4) controls the number of kernel fusion configurations XLA benchmarks during compilation, with level 4 running 50-200 autotuning benchmarks per compiled function, adding 30-120 seconds to compilation time but improving execution speed by 15-25%. The memory limit slop factor controls how aggressively XLA fuses operations despite memory pressure: lower values produce more conservative fusions that use less memory but launch more kernels. For latency-sensitive inference, `autotune_level=3` and `slop_factor=50` produce the best tradeoff for H100. For batch training throughput, `autotune_level=4` with `slop_factor=200` maximizes kernel fusion at the cost of 2-3 GB additional peak memory.

JAX Compilation SettingCompile TimeGPU MemoryInference LatencyThroughput Impact
autotune_level=1, slop=508-15 secondsBaseline +0 GBBaselineBaseline
autotune_level=2, slop=10015-30 seconds+1 GB-8%+5%
autotune_level=3, slop=5030-60 seconds+2 GB-15%+12%
autotune_level=4, slop=20060-180 seconds+3 GB-22%+18%
Default (autotune_level=2)15-30 secondsBaselineBaselineBaseline
Cache hit (subsequent runs)0.1-0.3 secondsSame as cachedSameSame
02

TPU VS GPU: JAX PERFORMANCE AND COST TRADEOFFS

JAX was developed by Google for TPU but has become the primary framework for GPU training at several major AI labs. The GPU backend uses the same XLA pipeline but with different lowering: TPUs use Pallas kernels on MXU (Matrix Unit) systolic arrays, while GPUs use Triton or CUDA kernels on tensor cores. The performance delta varies by operation: matrix multiplications (GEMM) run equivalently on TPU and GPU at comparable FLOPS budgets, but attention operations favor GPUs with FlashAttention kernels (JAX's `jax.nn.dot_product_attention` can call FlashAttention-2 via a custom VJP rule). Layer normalization and element-wise operations favor TPUs due to the MXU's higher bandwidth for fusion-friendly ops. A Llama 70B forward pass on TPU v5p (8x chips, 16 GB HBM each) completes in 38ms versus 32ms on 8x H100, making H100 16% faster for inference at similar cost.

The cost comparison favors GPUs for production deployment. TPU v5p pods cost approximately $4.50-6.00 per chip-hour on Google Cloud, versus H100 at $2.50-4.50 depending on provider. For a 256-GPU training cluster running 24/7, H100 is 35-50% cheaper than equivalent TPU compute. However, TPUs have two advantages: the PodSlice networking topology provides 95-98% inter-chip bandwidth utilization for all-reduce (versus 85-92% for GPU InfiniBand), and the XLA compilation for TPU is more mature with 40% faster compilation. The pragmatic pattern at most AI labs: develop in JAX on GPU for iteration speed and cost, perform full training runs on TPU pods for maximum throughput, then deploy inference on GPU for cost efficiency. This three-phase workflow is used by Google DeepMind, character.ai, and several foundation model labs.

DimensionH100 SXM 80 GBTPU v5pB200 SXMTPU v6 (Trillium)
Peak FP8 TFLOPS1,9792,7004,5004,100
HBM Capacity80 GB95 GB (per chip)288 GB192 GB
Interconnect BW900 GB/s NVLink4,800 GB/s ICI1,800 GB/s NVLink4,800 GB/s ICI
JAX Compile SpeedMedium (CUDA+HIP)Fast (mature XLA)Medium (CUDA)Fast (mature XLA)
Spot Price (est.)$2.35-3.50/hr$4.50-6.00/hr$3.90-4.50/hrN/A (2026+)
FlashAttention SupportYes (FA3)ExperimentalYes (FA3)Experimental
Ecosystem MaturityPyTorch dominantJAX nativePyTorch dominantJAX native
03

JIT COMPILATION CHALLENGES IN PRODUCTION

JAX's just-in-time compilation is the primary operational challenge for GPU deployment. The `jax.jit(f)` function traces and compiles on first invocation, which means the first batch of requests to a JAX-based inference service incurs a 30-180 second latency penalty while XLA compiles the model. For production serving, the standard workaround is ahead-of-time (AOT) compilation with `jax.jit(f).lower(x)` followed by `compile()` and `save()` to a serialized HLO file. The compiled artifact is loaded on server startup: `cached_fn = jax.load_compiled("/models/llama70b.hlo")`. This eliminates the JIT warmup latency but requires the compilation environment to match the serving environment's CUDA version and GPU architecture exactly. A mismatch causes XLA to recompile with a warning, reintroducing latency.

JAX recompilation on input shape change is another production issue. JIT traces functions based on `jax.ShapeDtypeStruct` arguments, so a model compiled for batch size 8 will recompile for batch size 9, double-caching both versions. For serving with variable batch sizes, `jax.jit(f, static_argnums=(0,))` marks batch size as a static argument to prevent recompilation, or `jax.jit(f, donate_argnums=(0,))` for buffer donation that reduces memory allocation overhead. The `jax._src.cache` module exposes cache statistics: `cache_info().misses` and `cache_info().hits`. A production JAX deployment should monitor the cache miss rate and alert if it exceeds 2% of total invocations, indicating excessive recompilation from shape variability. Host memory for the JAX cache also requires monitoring: each compiled function stores 50-500 MB of compiled PTX code, and a 64 GB cache for a multi-model serving deployment is not uncommon.

04

PJIT AND AUTOMATIC MODEL PARALLELISM FOR GPU CLUSTERS

JAX's `pjit` (partitioned JIT) is the distributed computation API equivalent to PyTorch's DTensor. It partitions computation across devices using `PartitionSpec` annotations that describe how each array dimension maps to a device mesh axis. For a Llama 70B model on 8x H100, the weight sharding is: `PartitionSpec('fsdp', None, None)` for the MLP weights (fsdp dimension = model sharding over devices), and `PartitionSpec('fsdp', None, None, None)` for attention weights with a batch dimension prefix. The `device_mesh = jax.make_mesh((8,), ('fsdp',))` creates a 1D mesh of 8 devices. `pjit` automatically inserts all-reduce collectives for any computation requiring cross-device data, using XLA's collective operation lowering rules.

JAX's `shard_map` offers explicit control over the sharding boundary, enabling manual optimization that `pjit`'s automatic approach cannot achieve. For expert parallelism in Mixture-of-Experts models, `shard_map` maps the expert computation over the device mesh dimension, with explicit all-to-all communication via `jax.lax.psum`. The DeepSeek-V3 MoE implementation in JAX uses `shard_map` for the expert routing, achieving 92% communication efficiency versus `pjit`'s 78% for the same topology. For standard transformer training, however, `pjit`'s automatic compilation is preferred for its simplicity. JAX's `jax.distributed.initialize()` replaces the older `jax.config.update("jax_xla_backend", "tpu_driver")` and must be called before any `jax.jit` or `pjit` operations in multi-process GPU configurations.

05

JAX INFERENCE SERVING ON GPU: REAL-WORLD PATTERNS

JAX inference serving on GPU requires a different infrastructure stack than PyTorch. The standard pattern uses `jax.jit` with `static_argnums` for fixed batch shapes and delivers compiled HLO artifacts through an HTTP server (typically FastAPI or JAX-native XLA:LA). The server precompiles the model into a `jaxlib.xla_extension.Executable` and calls `executable.execute(input_buffers)` in the request handler, achieving 2-3ms overhead per request outside model computation. For Llama 8B inference on H100, JAX achieves 380 tok/s at BS=1 versus 420 tok/s for PyTorch with `torch.compile`, a 10% deficit attributed to XLA's less aggressive CUDA graph reuse. For batch inference at BS=16, the gap narrows to 2-3% as XLA's fusion optimizations become more effective at higher arithmetic intensity.

The choice between JAX and PyTorch for GPU deployment hinges on the training framework. If the model was trained in JAX (common at Google, character.ai, and several foundation labs), deploying in JAX avoids weight conversion and numerical precision mismatches. The Safetensors weight export format with `jax.numpy.save` and `flax.serialization` enables cross-framework weight sharing, but the computational graphs differ between JAX and PyTorch, making identical inference output difficult to guarantee. On ClusterBid, JAX-based GPU serving is most common at AI labs with JAX training pipelines, while PyTorch dominates for teams using Hugging Face or Megatron-LM. The GPU infrastructure requirements are identical: the same H100, A100, and B200 instances serve both frameworks with equivalent performance for compute-bound workloads.

Filed under
JAX GPUXLA CompilationJAX TPU vs GPUJIT Compilation JAXpjit ShardingJAX H100 InferenceJAX Production Deployment