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.

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.

Measured inputs and delivered differences
Measure / unitBeforeAfterDifference
Monthly accepted-job compute charge
USD/month
12,0008,0004,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

Which layer should change before the company commits?

Select a step to follow its reasoning, then continue into the technical chapters.

Problem → boundary → decision → evidence

Separate the layers

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.
Normalize hardware · Decision
Keep precision, memory and interconnect scope comparable.
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.

python / example
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

From HBM bytes to a result inside a GPU

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

HBM3e DRAM

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.

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 assumptions

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
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()

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

Engineering notes

Keep following the thread.

Real client engagements and the engineering behind them.

Technical glossary: definitions, connected ideas and further reading.

Optional analytics off. Contact works either way.

How measurement works