Hover a dotted term for 5 seconds to lock its explanation. It closes after 5 seconds away; nested tooltips and keyboard focus keep it open. Click, tap or Enter locks immediately. Technical glossary.

Compilers & accelerators / Interactive field guide

JAX, CUDA, TPU and Blackwell: follow the execution boundary

Follow the programming model into the compiler and device. The useful comparison is who controls each boundary—not which unrelated peak number wins a chart.

4 linked boundariesLocal teaching modelNo account required

01 / Follow the explanation

The system, step by step.

A brief visual sequence plays automatically. The example and its assumptions are already here—nothing to configure.

An interactive teaching model drawn from real delivery work. No agent, cloud account, GPU or cluster is accessed.

Execution responsibilities: program frontend, compiler, device execution, multi-device collective communication.A responsibility map, not a guaranteed lowering pipeline. Compilers may choose libraries, generated kernels or custom calls. Collective work is relevant only for distributed execution.01Program02Compiler03Device04Collective
A responsibility map, not a guaranteed lowering pipeline. Compilers may choose libraries, generated kernels or custom calls. Collective work is relevant only for distributed execution.

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

Programming and hardware path
JAX → NVIDIA GPU / B200
Programming model
JAX array transformations
Execution path
Tracing to compiler IR to executable

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.

Compiler

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.

Collective

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.

Instructional excerpt, not executed here
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

Transparent reasoning / Units and boundaries
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.

  1. Which JAX transformations remain when moving from GPU to TPU, and which layout and collective behavior changes?

  2. Why would timing host dispatch without waiting for device completion answer the wrong question?

  3. Why does explicit CUDA kernel control not automatically provide multi-GPU sharding or eliminate communication?

Primary documentation

References & further reading

Engineering notes

Read the project behind the model.

Real client engagements and the engineering behind them.

A useful next conversation

What needs to work better?

A system, a delivery bottleneck, or an engineering opportunity. Tell me what you are building and where you want to go.

Let’s talk

Technical glossary: definitions, connected ideas and further reading.

Optional analytics off. Contact works either way.

How measurement works