Current example
JAX on NVIDIA GPU: trace, lower, then compile
JAX traces function execution with abstract values such as shape and dtype, not the Python AST. It lowers a specialized computation to compiler IR and compiles for the selected NVIDIA GPU backend. Compilation and caches depend on signatures and configuration; JAX can run on NVIDIA GPUs as well as TPUs. This is a pipeline explanation, not a hardware speed score.
Example path: JAX on an NVIDIA GPU, compilation stage. TPU means v6e (Trillium); Blackwell means B200 in an HGX context, not B300, GB200 NVL72 or an entire accelerator rack.
Choose the path
Inspect a stage
- Execution stage
- Trace and compile
- Boundary to inspect
- Shape, dtype and backend specialization
02 / Follow the flow
Four boundaries. One connected explanation.
Each card explains a node in the diagram. Highlights show which boundaries contribute to the current result.
Program
Choose the level where the algorithm lives
JAX exposes array operations and transformations such as jit, grad and vmap. Tracing executes a Python function with abstract values to build a program representation; it is not general compilation of arbitrary Python source. Direct CUDA exposes thread indexing, blocks, shared memory and synchronization. These layers can coexist: a high-level JAX program may call optimized GPU libraries or a custom kernel rather than replacing them.
Separate specialization from steady execution
JAX lowering and backend compilation depend on shapes, dtypes and static arguments. A changed signature can require another compilation. CUDA C++ is compiled with an explicit target architecture, with possible driver-side compilation depending on the artifact. Neither a jit annotation nor a successful compiler invocation proves that the intended fusion, layout or Tensor Core path was selected. Inspect compiler output and profiler evidence for the exact toolchain.
Device
Measure completed work on an exact configuration
Device dispatch is asynchronous. JAX timing needs a completion boundary such as block_until_ready; CUDA event timing needs the relevant stream and event completion. State whether inputs already reside on the device and whether compilation, transfers and host work are included. TPU v6e uses matrix, vector and scalar execution units; B200 is a Blackwell GPU with its own instruction and memory hierarchy. A framework label alone does not determine their performance.
Scaling changes the data-movement problem
Sharding is an algorithmic choice about where values live and which reductions or exchanges are required. TPU v6e uses an ICI fabric with a 2D torus topology at Pod scale. HGX B200 connects eight GPUs through a fifth-generation NVLink/NVSwitch scale-up fabric; host networking is a different scale-out boundary. A TPU Pod is an accelerator system, not a Kubernetes Pod. Compare equivalent model partitions, communication volumes and failure boundaries—not a single GPU against an entire TPU slice.
import jax
import numpy as np
# Requires a configured GPU backend; no CPU fallback.
device = jax.devices("gpu")[0]
x = jax.device_put(np.ones((256, 256), np.float32), device)
f = jax.jit(lambda a: a @ a)
lowered = f.lower(x) # trace abstract shape/dtype, then lower
print(lowered.as_text()) # inspect compiler IR, not assembly
compiled = lowered.compile()
y = compiled(x)
y.block_until_ready()03 / Keep the model honest
Execution assumptions
Comparison boundary = workload + shapes + dtype + quality target
Cold path = initialization + transfers + compilation + execution
Warm device path = prepared inputs + completed device work
Distributed path adds sharding + collectives + synchronization
No throughput, latency, cost or energy winner is calculated here.- JAX is a numerical programming and transformation system; CUDA is a GPU platform and programming model. Selecting JAX does not imply selecting TPU.
- TPU v6e and Blackwell B200 are named reference architectures, not claims about the newest available generation or every cloud offering.
- Published chip, GPU, eight-GPU baseboard and rack specifications are different units. Preserve decimal GB versus binary GiB and dense versus sparse precision qualifiers.
- Do not compare a sparse low-precision peak with dense higher-precision execution. Kernel support, numerical quality, memory fit and achieved utilization need workload-specific evidence.
- The snippets are code readings, not programs executed by this browser. There is no GPU runtime, compiler download or hardware benchmark.
- All four responsibility explanations and the worked example remain available without JavaScript.
04 / Think it through
Questions behind the example.
Which JAX transformations remain when moving from GPU to TPU, and which layout and collective behavior changes?
Why would timing host dispatch without waiting for device completion answer the wrong question?
Why does explicit CUDA kernel control not automatically provide multi-GPU sharding or eliminate communication?
Primary documentation

