Accelerator engineering
JAX vs CUDA and GPU vs TPU: A Workload-Based Decision Guide
A real client engagement. The engineering and the results are described below.
A forecasting company we worked with is considering a platform move to meet a larger customer’s overnight deadline. The team separates JAX, CUDA and accelerator hardware into different decisions before authorizing a rewrite.
Request a time through the inquiry form. A meeting is confirmed separately by email.
The business problem behind the technology
Which layer should change before the company commits?
Every candidate must retain model quality, supported shapes and the delivery deadline. A short warm kernel timing cannot stand in for compilation, transfer and the complete job.
Read the client engagement ↓Client engagement / Delivered results
A platform choice became a reversible workload experiment
A forecasting provider we worked with runs 100 scheduled training-and-evaluation jobs a month. Its leadership sees cheaper advertised compute and asks whether a different framework or accelerator would free budget.
The constraint
Every candidate must retain model quality, supported shapes and the delivery deadline. A short warm kernel timing cannot stand in for compilation, transfer and the complete job.
The engineering decision
The team tests exact configurations and records cold/warm paths separately. In this engagement, a baseline job took ten billable hours at $12/hour; the candidate took eight at $10/hour, with identical accepted output.
The delivered outcome
The monthly job compute charge fell from $12,000 to $8,000. That difference supported a pilot, not a general conclusion that a TPU, GPU or framework is cheaper.
| Measure / unit | Before | After | Difference |
|---|---|---|---|
| Monthly accepted-job compute charge USD/month | 12,000 | 8,000 | 4,000 |
Monthly accepted-job compute charge. 100 × 10 × $12 versus 100 × 8 × $10 for this workload.
The conditions behind the results
- The 100 jobs use equivalent data, quality, numerical precision and completion criteria.
- Billable durations include the relevant setup and transfer scope, not only warm kernel execution.
- Porting, staff, storage, networking and long-term commitments are excluded from this narrow ledger.
What this does not prove. This compute difference is specific to the client’s jobs; it does not rank vendors or accelerators universally.
Evidence to collect for your own decision
- Separate tracing, lowering, compilation, transfer and synchronized execution in the evidence.
- Compare the exact device count, topology, precision and quality acceptance.
- Price the migration and fallback path before making an irreversible platform commitment.
Key decisions
JAX and CUDA are not opposing hardware choices. Follow the program through tracing, lowering, device execution, and collectives; then compare a precisely defined workload on precisely named accelerator configurations.
- Separate tracing, lowering, compilation, transfer and synchronized execution in the evidence.
- Compare the exact device count, topology, precision and quality acceptance.
- This compute difference is a workload-model comparison, not a vendor ranking, accelerator benchmark or full migration ROI.
Follow the decision
Select a step to follow its reasoning, then continue into the technical chapters.
Problem → boundary → decision → evidence
Distinguish a framework, a programming interface and a device.
Read every component and connection
- Separate the layers · Problem
- Distinguish a framework, a programming interface and a device.
- Inspect compilation · Boundary
- Trace lowering and the implementation actually selected.
- Choose the next test · Evidence
- Use synchronized workload evidence to justify the change.
- Separate the layers → Inspect compilation: identify the constraint
- Inspect compilation → Normalize hardware: choose a bounded change
- Normalize hardware → Choose the next test: check the outcome
A conceptual decision map for this article, not a measured timeline, physical topology or a depiction of a specific client system.
Hover, focus or tap a component to inspect it. Motion adapts automatically to connection, device and accessibility signals; the component key remains readable without JavaScript.
Put the choices at the correct abstraction layers
The client’s forecasting team separates framework and accelerator decisions before pricing a migration.
JAX is a Python array-programming and transformation system: it can compose automatic differentiation, vectorization, compilation, and sharding around numerical functions. CUDA is an NVIDIA programming platform with a runtime, libraries, and a kernel execution model. Writing CUDA C++ gives direct responsibility for launches, thread/block organization, memory access, and synchronization. Comparing JAX with CUDA as though they were two accelerator chips mixes levels of the stack.
JAX can run on NVIDIA GPUs through its GPU backend, which can ultimately execute CUDA kernels and CUDA libraries. JAX can also target TPUs through their backend. The actual choice might be JAX on B200 versus JAX on TPU v6e, or a JAX-generated GPU operation versus a hand-written CUDA operation on the same B200. Those are different experiments with different controlled variables. CUDA blocks are groups of device threads, not Kubernetes pods; a Cloud TPU Pod likewise names a hardware system, not a Kubernetes pod.
JAX traces Python execution, not the Python AST
The Python function becomes a traced computation. You inspect what that means before attributing a result to the device.
Under ordinary jax.jit tracing, Python executes the function with tracer values that carry abstract properties such as shape and dtype. JAX records the JAX primitive operations performed during that execution into a jaxpr. It does not parse a Python abstract syntax tree and automatically translate arbitrary Python semantics into accelerator code. The distinction explains both its transformation power and its sharp edges.
A Python branch selected from a static property, such as array rank, can be decided during tracing. A Python if that requires the runtime numerical value of a traced array generally cannot be decided that way; express supported data-dependent control flow with JAX control-flow primitives such as jax.lax.cond. Ordinary Python side effects are not automatically represented in the compiled computation. A print or mutation observed while tracing is not proof it will happen on every execution.
Compiled artifacts specialize to relevant input signatures and static arguments. A new shape, dtype, or static value can require another trace and compilation. Avoid recreating jitted functions inside a hot loop and track the actual workload shape distribution. Shape bucketing may reduce specialization churn but can introduce padding and extra compute. Measure that tradeoff rather than attributing every first-call delay to a slow accelerator.
Lowering chooses libraries, generated kernels, and boundaries
Lowering selects libraries, generated kernels and boundaries. The source expression alone does not name the executed implementation.
JAX lowers its staged computation toward StableHLO and the XLA backend pipeline. OpenXLA describes shape/layout decisions, sharding, fusion, buffer assignment, and scheduling as parts of that process. On GPU, an operation may become a library call, compiler-emitted code, or a generated fused kernel. Source-level array expressions do not imply one kernel per expression, and a whole jitted function does not imply exactly one kernel either.
The XLA:GPU architecture documentation names cuBLAS, cuDNN, and NCCL as library paths and native/PTX and Triton-based code generation as other choices. Which implementation is selected depends on the operation, shape, dtype, backend, compiler version, and configuration. Fusion can avoid materializing intermediate arrays in high-bandwidth memory, but register pressure, layout conversions, and lost library opportunities can change the outcome. Inspect optimized compiler output and an execution profile before prescribing a rewrite.
Use the maintained jitted-function lowering API, such as jitted_function.lower(inputs).compile(), when inspecting the staged boundary; do not assume every historical documentation example uses a current API. An optimized graph explains compiler decisions, while device profiling explains time spent. Neither replaces numerical correctness checks against a reference and declared tolerance.
Separate lowering and compilation from synchronized warm calls
The first call includes work a warm call may not. Compilation and synchronization need separate treatment.
JAX dispatch is asynchronous: receiving an array object is not necessarily evidence that its device work has finished. A host timer around an unsynchronized call can primarily measure dispatch. For a warm-call latency sample, finish input placement first, lower and compile outside the sample, run a separate warm-up, then stop the host timer only after the result is ready. block_until_ready waits for the result without requiring a host copy of its values.
The excerpt below uses float32 inputs, a fixed 512 by 512 matrix shape, the selected default backend, and a single returned array. Its lower-and-compile interval includes tracing/lowering and compilation work that occurs in that call; existing caches can affect it, so it is not guaranteed to be a cold compile. The warm samples include Python invocation, runtime dispatch, execution, and synchronization, but exclude initial input placement and copying result data back to the host. They are not pure kernel times.
Run latency and throughput experiments separately. Blocking after every invocation serializes the samples and measures this call pattern; an asynchronous production pipeline can overlap host work and device execution differently. Record device identity, JAX/jaxlib and backend versions, compilation-cache state, shape, dtype, precision settings, and sample count. The article provides no timing output and makes no GPU-execution claim.
import time
import statistics
import numpy as np
import jax
@jax.jit
def matmul(a, b):
return a @ b
rng = np.random.default_rng(0)
a = jax.device_put(rng.standard_normal((512, 512)).astype(np.float32))
b = jax.device_put(rng.standard_normal((512, 512)).astype(np.float32))
a.block_until_ready()
b.block_until_ready()
t0 = time.perf_counter()
compiled = matmul.lower(a, b).compile()
lower_compile_s = time.perf_counter() - t0
compiled(a, b).block_until_ready() # separate execution warm-up
samples_s = []
for _ in range(20):
t0 = time.perf_counter()
compiled(a, b).block_until_ready()
samples_s.append(time.perf_counter() - t0)
print("devices:", jax.devices())
print("lower + compile seconds:", lower_compile_s)
print("warm median milliseconds:", 1000 * statistics.median(samples_s))CUDA events measure a different timing boundary
CUDA events offer another timing boundary. You keep their meaning distinct from host-side request time.
CUDA kernel launches enqueue work into streams. A stream orders its submitted operations; independence across streams does not guarantee overlap, and dependencies or hardware resources can prevent it. A CPU timestamp immediately after a launch does not delimit kernel completion. Check both launch errors and execution errors surfaced at synchronization rather than accepting a timing from a failed operation.
For an isolated device interval, record a start CUDA event and a stop CUDA event in the same stream around the work, synchronize the stop event, and obtain milliseconds with cudaEventElapsedTime. Create timing events outside the measured interval and destroy them when done. NVIDIA documents that event timestamps are recorded when the device reaches those events, so this differs from host end-to-end timing. Other queued work, concurrent streams, and dependencies can still affect the observed interval.
For a multi-stream pipeline, explicitly establish event dependencies and a completion boundary covering all relevant work; two events in an unrelated stream do not enclose it. Decide whether host-to-device and device-to-host transfers belong in the comparison. Compare a JAX synchronized host interval with an equivalent CUDA host interval, or profile equivalent device intervals on both paths. Do not label a dispatch-only measurement versus an event-delimited kernel measurement as an implementation speedup.
Name TPU v6e chips and B200 systems precisely
The exact chip, node and communication boundary must match the job-cost estimate.
Google identifies TPU v6e as Trillium. Its architecture page describes one TensorCore per chip containing two matrix-multiply units, a vector unit, and a scalar unit. The published per-chip HBM capacity is 32 GB, while a full v6e Pod contains 256 chips in a two-dimensional torus. A selected slice or VM is a smaller explicitly configured unit. Do not compare a TPU chip with an eight-GPU server and call both one accelerator.
The linked NVIDIA Blackwell datasheet identifies the HGX B200 platform as eight B200 GPUs and lists 180 GB of HBM3E per GPU, with 1.4 TB as the rounded platform memory specification. B200 is the named GPU here; HGX B200 is the multi-GPU baseboard/platform. Neither is interchangeable with GB200 NVL72, a different rack-scale configuration, or B300, a different Blackwell offering. The datasheet describes Blackwell GPUs as two dies connected into one logical GPU, not two separately schedulable GPU resources.
Retain the vendors’ published GB and TB labels rather than silently rewriting them as GiB and TiB. In arithmetic that requires byte conversion, state the convention: decimal GB = 10^9 bytes and binary GiB = 2^30 bytes. Verify the actual device-reported usable capacity for a deployment; advertised memory is not all available for model weights. Aggregate platform memory also does not become a free, uniform allocation pool: a model must shard or otherwise manage data across devices.
Conceptual Blackwell execution and memory map
Off-chip DRAM → memory controllers → caches → registers / execution → stores. Instruction issue and asynchronous copies coordinate different paths.
Large off-chip storage
On-chip reuse and staging
Inside one representative streaming multiprocessor
Device-wide work and peers
Stacked DRAM stores large device-resident arrays. DRAM cells need refresh; capacity and sustained bandwidth are different limits. HBM is not the register file.
Read every component and connection
- HBM3e DRAM · Weights / KV / arrays
- Stacked DRAM stores large device-resident arrays. DRAM cells need refresh; capacity and sustained bandwidth are different limits. HBM is not the register file.
- Memory controllers · Channels and requests
- Controllers organize reads and writes to memory channels. Access patterns, contention and the memory technology influence service time; bandwidth is not zero latency.
- L2 cache · Device-wide reuse
- L2 can satisfy repeated requests without another HBM access. Its capacity and residency behavior affect traffic; a cache hit is not a new DRAM transfer.
- L1 / shared memory · Caching / explicit tiles
- B200 combines L1, texture and shared-memory resources. Shared memory is software-managed block storage with synchronization rules; it is not an automatic replacement for registers.
- Register file · Thread operands
- The B200 tuning guide specifies 64K 32-bit registers per SM. Threads use registers for live values; spills can create device-memory traffic. Registers are not off-chip DRAM.
- ALU pipelines · Integer / floating point
- Execution pipelines perform supported arithmetic and logic on operands. CMOS gates underlie those circuits. Floating-point operations include more work than the integer full adder shown.
- Tensor cores · Matrix operations
- Specialized matrix instructions use supported operand formats and accumulation paths. Tensor throughput is not scalar ALU throughput, and not every kernel can use tensor cores.
- Load / store units · Addresses and movement
- Load/store machinery forms and services memory operations. Coalescing groups useful lane accesses; dependencies prevent a consumer from using a value before it is ready.
- Warp schedulers · Ready instruction issue
- Schedulers issue eligible warp instructions subject to dependencies and resource availability. Other ready warps can hide a wait; occupancy alone does not prove throughput.
- Instruction path · Fetch / decode / issue
- Compiled machine instructions reach the SM instruction machinery. PTX is a virtual ISA; a compatible cubin or driver compilation supplies hardware-executable code.
- Async copy / TMA · Tile movement
- Supported asynchronous transfer paths can stage tiles while computation proceeds. Barriers and producer/consumer ordering still apply; overlap is not permission to read unfinished data.
- GPU front end · Submitted work
- Device work submission and scheduling machinery distribute kernel work. Block resource requirements influence residency. Kubernetes does not choose a warp or allocate an SM register.
- NVLink interface · Peer devices
- Peer access and collectives move data between compatible GPUs. The application/runtime manages distributed work; aggregate device memory is not one automatically shared allocation.
- GPU front end → Instruction path: kernel work
- Instruction path → Warp schedulers: decoded instructions
- Warp schedulers → Load / store units: memory instruction
- HBM3e DRAM → Memory controllers: DRAM service
- Memory controllers → L2 cache: cache-line traffic
- L2 cache → L1 / shared memory: cache / tile path
- L1 / shared memory → Load / store units: load path
- Load / store units → Register file: operand load
- Register file → ALU pipelines: ALU operands
- ALU pipelines → Register file: result
- Register file → Load / store units: store
- L2 cache → Async copy / TMA: async tile copy
- Async copy / TMA → L1 / shared memory: staging
- L1 / shared memory → Tensor cores: supported matrix operands
- L2 cache → NVLink interface: peer traffic
This is a functional map, not a floorplan or cycle-accurate simulator. One representative SM is expanded; it is not the GPU’s SM count. Cache bypass, asynchronous copies, distributed shared memory and specialized tensor accumulator paths mean not every operation follows every arrow.
Hover, focus or tap a component to inspect it. Motion adapts automatically to connection, device and accessibility signals; the component key remains readable without JavaScript.
Interconnect figures need direction and topology
Several devices add communication. Direction, topology and usable paths matter more than an isolated bandwidth headline.
The v6e architecture page lists 1,638 GB/s of HBM bandwidth and 800 GB/s of bidirectional inter-chip interconnect bandwidth per chip, with four ICI ports. These are different interfaces: on-chip-attached memory traffic is not cross-chip traffic. A torus slice shape changes communication paths, and a whole-Pod aggregate or bisection figure measures something different from one chip’s interface.
For B200, the Blackwell datasheet lists 7.7 TB/s of HBM bandwidth per GPU. The HGX page distinguishes 1.8 TB/s of GPU-to-GPU NVLink bandwidth from 14.4 TB/s of total NVLink bandwidth for the eight-GPU system. These are vendor interface specifications, not a payload rate for a chosen all-reduce. Preserve their aggregation and direction conventions instead of multiplying another system’s one-way number into an apparent advantage.
The host and cross-host networks remain additional boundaries. An eight-GPU NVLink-connected group does not promise NVLink performance between arbitrary servers, and ICI is not the same link as a TPU host’s data-center NIC. Compare the actual collective topology, message sizes, contention, and effective payload bandwidth at the proposed scale. No single advertised bandwidth number determines distributed step time.
Precision and shape determine what a hardware comparison means
The numerical contract narrows the candidate implementations. Precision and shape can change which hardware path is available.
A floating-point label needs context: input storage dtype, multiply precision, accumulator precision, output dtype, and numerical tolerance can differ. A float32 array expression is not by itself a guarantee of a particular accelerator matrix-instruction mode. BF16, FP16, FP8, and FP4 are not interchangeable accuracy contracts. Quantized execution also includes scaling, conversion, supported operation coverage, and validation against the application task.
This article deliberately omits peak-compute rankings. NVIDIA’s HGX and Blackwell tables distinguish sparse from dense Tensor Core specifications; Google’s v6e page labels its per-chip compute rows by dtype. A defensible table would have to align exact dtype and arithmetic mode, dense versus supported structured-sparse work, operation counting, and per-chip versus whole-system scope before any ratio is meaningful. A sparse peak is not available to an arbitrary dense model, and a vendor peak is not a benchmark result.
Large regular matrix workloads, small elementwise chains, irregular gathers, long-context attention, and dynamic shapes stress different parts of the machine. Tiling and alignment can favor some shapes; padding and reshaping have costs; small jobs may spend more time in dispatch or transfer than arithmetic. Build a workload matrix from real intended shapes and data movement. The right question is which configuration meets correctness, latency, throughput, and cost requirements, not which logo wins a utilization chart.
Portability and control can coexist at a narrow boundary
Portability does not require surrendering every low-level choice. A narrow boundary can preserve control where the evidence justifies it.
Start with a clear numerical function and a reference implementation. JAX can preserve a common high-level program across CPU, GPU, and TPU backends when its operations and dependencies are supported, but source portability is not automatic performance portability. Backend-specific layouts, sharding, precision behavior, and operation availability still need validation. Compilation and runtime versions are part of the experiment.
Direct CUDA is useful when a measured hotspot needs a carefully controlled memory layout, synchronization pattern, or hardware-specific operation that the current compiled path does not provide well. It also makes the author responsible for bounds, races, numerical behavior, launch choices, and long-term compatibility. Compare against the relevant optimized library, not only a deliberately naive baseline, before accepting that maintenance cost.
JAX’s documented foreign-function interface can call existing compiled C or CUDA implementations, so a custom kernel does not require abandoning the entire JAX program. The boundary must declare shapes and dtypes and integrate correctly with runtime streams and buffer ownership. JAX does not automatically know how to differentiate a foreign function; transformation and sharding behavior must be supplied or constrained as appropriate. Keep such a boundary narrow and evidence-driven, rather than treating FFI as a universal optimization switch.
Multi-device execution is a data-placement problem
The workload spreads across devices, and placement becomes part of the algorithm’s practical cost.
A device mesh and sharding specification describe where logical array pieces reside. XLA’s partitioning can insert collectives and resharding to satisfy the program, but a concise array expression does not erase communication. Inspect whether an intermediate is replicated, partitioned along a useful dimension, or repeatedly gathered and redistributed. Capacity fits must include temporary communication buffers as well as the final shards.
Data parallelism replicates model work across different examples and requires the appropriate synchronization for the algorithm. Tensor/model sharding changes where parts of an operation or model live and can require collectives inside a step. Measure both a fixed global workload when scaling devices and any changed per-device batch size. Otherwise an apparent scaling improvement may just be a different problem size.
Account for input staging, host preprocessing, device-to-device transfers, collective synchronization, and output consumption. Keep intermediate arrays on devices when possible; pulling them back to Python can introduce transfers and synchronization. Record the slowest participating worker and the completion boundary for the full step. A faster isolated matmul can leave end-to-end performance unchanged when the critical path is an input pipeline or a collective.
Use the execution reader to choose the next measurement
The next measurement determines whether the compute difference survives setup, quality and operating costs.
The companion reader switches among JAX on GPU, JAX on TPU, and direct CUDA, then highlights compilation, execution, or scaling. It is a map of responsibility and timing boundaries, not a simulator of a physical accelerator. Changing the selected backend does not run code or manufacture a speed estimate. Use the highlighted stage to ask a concrete question: are we retracing, moving data, waiting on device work, or communicating between devices?
For a comparison, keep the model or numerical operation, representative inputs, output checks, precision policy, and measurement boundary fixed. Report lowering/compilation separately from warmed execution, then include a complete application path with transfers and output consumption. Repeat at realistic shapes and scales, state cache state, and retain sample counts and variability. If a custom CUDA kernel helps only one narrow shape, say so rather than extending the claim to JAX as a whole.
The decision record should explain what bottleneck was observed, which layer changed, whether correctness remained within the declared tolerance, and whether the result still holds at the deployment scale and service objective. Include engineering and operating cost alongside performance. Until those observations exist, the honest conclusion is a testable hypothesis about workload fit, not an unsourced TPU-versus-GPU winner.
01 / Follow the explanation
JAX, CUDA, TPU and Blackwell: follow the execution boundary
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.
Read the full explanation and assumptionsCurrent 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
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()Questions behind the decision
Are JAX and CUDA competing hardware choices?
No. JAX is a high-level numerical framework with compilation paths; CUDA is a programming platform for supported NVIDIA GPUs. Compare software control and hardware selection separately, then measure the actual workload.
How should GPU and TPU costs be compared?
Use the same accepted model quality and complete job boundary on precisely named configurations. Include compilation, data placement, communication and commercial commitments before translating runtime into cost.
References & further reading
- JAX: tracing, jaxpr, JIT specialization, and synchronized warm-up
- OpenXLA: GPU lowering, fusion, partitioning, and library selection
- JAX: foreign-function interface and custom differentiation boundaries
- NVIDIA: CUDA timing, data movement, and correctness guidance
- Google Cloud: TPU v6e per-chip architecture and slice configurations
- NVIDIA: HGX platform versus individual GPU and interconnect specifications
- NVIDIA: Blackwell datasheet, B200 memory, system units, and precision footnotes


