Framework comparisons usually list features: this one has a model zoo, that one has a nicer API. Those lists go stale quickly and rarely help you choose. What does not go stale is the execution model, meaning how Python code turns into work on a GPU. PyTorch runs operations eagerly as your Python executes and optionally compiles; JAX asks you to write pure functions that it traces and compiles as a whole; TensorFlow traces Python into a graph and can hand that graph to the XLA compiler. Almost every practical difference, from debugging to multi-GPU scaling to the bugs you will hit, follows from that one choice.
This article explains each model from first principles, writes the same training step in all three, and follows each down to the kernels and collectives it launches. It assumes you know what a gradient is, and nothing about framework internals.
The one distinction that matters: eager versus traced
In an eager framework, each line such as y = x @ w immediately asks the GPU to do the work. Python sits in the loop: the interpreter dispatches the operation, the framework picks a kernel, enqueues it on a CUDA stream and returns a tensor handle while the GPU runs asynchronously. You can print intermediate values, use ordinary if statements on data and step through in a debugger. The cost is overhead per operation, a few microseconds of CPU time each, and the fact that the framework only ever sees one operation at a time, so it cannot fuse neighbours into one kernel.
In a traced framework, the first call runs your Python with placeholder values that record operations instead of computing them. The recording, a graph of the whole function, goes to a compiler that can fuse element-wise operations into matrix multiplies, plan memory, and choose layouts. Later calls with the same input shapes and types skip Python entirely and launch the compiled program. The costs are compile time, re-tracing whenever shapes change, and the rule that Python side effects run only while tracing, not on every call.
PyTorch defaults to eager and adds tracing through torch.compile. JAX is traced by design, and its functional style exists to make tracing safe. TensorFlow 2 is eager by default, but tf.function traces Keras training loops into graphs. The real question is how much of your program lives inside a compiled region, and how painful it is to get it there.
One training step, written three ways
The same two-layer network, AdamW and bf16 compute, in each framework. Read them for what they reveal about state, not for syntax.
import torch
import torch.nn as nn
import torch.nn.functional as F
model = nn.Sequential(nn.Linear(1024, 4096), nn.GELU(), nn.Linear(4096, 10)).cuda()
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)
compiled = torch.compile(model) # optional: capture + generate fused kernels
def train_step(x, y):
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = F.cross_entropy(compiled(x), y)
loss.backward() # autograd recorded the graph during forward
opt.step() # mutates parameters in place
opt.zero_grad(set_to_none=True)
return loss.detach()PyTorch is object-oriented and mutable. Parameters live inside the module, backward() writes gradients into .grad fields, and the optimizer updates tensors in place. Autocast chooses bf16 for matrix multiplies and keeps reductions in float32. torch.compile wraps the module; if it cannot capture something, it falls back to eager for that piece rather than failing.
import jax
import jax.numpy as jnp
import optax
def init(key):
k1, k2 = jax.random.split(key)
return {"w1": jax.random.normal(k1, (1024, 4096)) * 0.02, "b1": jnp.zeros(4096),
"w2": jax.random.normal(k2, (4096, 10)) * 0.02, "b2": jnp.zeros(10)}
def loss_fn(params, x, y):
h = jax.nn.gelu(x @ params["w1"] + params["b1"])
logits = h @ params["w2"] + params["b2"]
return optax.softmax_cross_entropy_with_integer_labels(logits, y).mean()
opt = optax.adamw(3e-4)
@jax.jit
def train_step(params, opt_state, x, y):
loss, grads = jax.value_and_grad(loss_fn)(params, x, y)
updates, opt_state = opt.update(grads, opt_state, params)
return optax.apply_updates(params, updates), opt_state, loss # new values, no mutation
params = init(jax.random.key(0))
opt_state = opt.init(params)JAX is functional. Parameters are a plain nested dictionary (a pytree), jax.value_and_grad transforms the loss function into one that also returns gradients, and the step returns new parameters instead of modifying old ones. Randomness takes an explicit key. Nothing is hidden in objects, so jit can trace the whole step, backward pass and optimizer included, into one XLA program. Flax (including its newer NNX API) and Equinox add module ergonomics on top.
import tensorflow as tf
model = tf.keras.Sequential([tf.keras.layers.Dense(4096, activation="gelu"),
tf.keras.layers.Dense(10)])
opt = tf.keras.optimizers.AdamW(3e-4)
loss_obj = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
@tf.function(jit_compile=True) # trace once per input signature, compile with XLA
def train_step(x, y):
with tf.GradientTape() as tape:
loss = loss_obj(y, model(x, training=True))
grads = tape.gradient(loss, model.trainable_variables)
opt.apply_gradients(zip(grads, model.trainable_variables))
return lossTensorFlow sits between them. Keras layers own variables as PyTorch modules do, GradientTape records operations for differentiation, and tf.function traces the step. With jit_compile=True the traced graph is compiled by XLA, the same compiler JAX uses; without it, the graph runs on TensorFlow's own kernel runtime with graph-level optimisation. Keras 3 is multi-backend and can run the same model code on JAX, TensorFlow or PyTorch, which is useful if your team writes Keras but wants a different engine underneath.
Down to the GPU: dispatch, compilers and kernels
The bottom two rows are nearly identical across frameworks. A large bf16 matrix multiply ends up in cuBLAS or a similar vendor kernel running on tensor cores regardless of which Python API asked for it, which is why framework choice rarely changes the speed of a single big GEMM. See how tensor cores execute those multiplies for that layer.
The difference is in everything between the multiplies: the bias adds, activations, normalisations, dropout masks and optimizer arithmetic. Eager PyTorch launches a kernel for each and round-trips every intermediate through GPU memory. A compiler sees them together and fuses chains of element-wise work into single kernels that read and write memory once. For memory-bound models, which includes most transformer training outside the big matrix multiplies, fusion is where compiled code wins; the kernel fusion article explains why bandwidth, not arithmetic, is the limit.
PyTorch's Inductor generates Triton kernels for fusion and calls vendor libraries for matrix multiplies. XLA, used by JAX and optionally TensorFlow, compiles the whole program, including buffer assignment. XLA sees more at once; Inductor tolerates messier Python because Dynamo can fall back to eager.
Compilation in practice: what breaks and what it costs
Every tracing system shares three hazards. Shape changes trigger recompiles. A traced program is specialised to input shapes, so a data loader that emits variable sequence lengths can compile a new program per batch. Pad or bucket to a small set of shapes. Data-dependent Python control flow cannot be traced as Python. An if loss > threshold on a tensor value is unknown at trace time; JAX raises a concretisation error and asks you to use jax.lax.cond or jnp.where, TensorFlow converts some control flow automatically through AutoGraph, and torch.compile inserts a graph break and runs that part eagerly. Side effects run at trace time. A print or a Python list append inside a jitted function runs once, during tracing, not every step.
PyTorch degrades gracefully: code always runs, but graph breaks can quietly leave half your model eager. TORCH_LOGS=graph_breaks lists them, and fullgraph=True turns a break into an error. JAX is strict: if it compiles, the whole step is one program, but you must write in its style from the start. TensorFlow retracing is controlled by input_signature, and unexpected retracing is a common cause of slow Keras jobs.
Compile time is real: minutes for a large model on the first step, and mid-run recompiles look like mysterious stalls. Measure warm-up separately from steady state.
Scaling out: who inserts the collectives
Multi-GPU training needs gradient all-reduce for data parallelism, and all-gather and reduce-scatter when parameters are sharded. The frameworks differ in who decides where those collectives go.
# JAX: describe WHERE arrays live; the compiler inserts the collectives.
from jax.sharding import NamedSharding, PartitionSpec as P
mesh = jax.make_mesh((jax.device_count(),), ("data",))
params = jax.device_put(params, NamedSharding(mesh, P())) # replicated
opt_state = jax.device_put(opt_state, NamedSharding(mesh, P()))
x = jax.device_put(x, NamedSharding(mesh, P("data"))) # batch split across GPUs
y = jax.device_put(y, NamedSharding(mesh, P("data")))
params, opt_state, loss = train_step(params, opt_state, x, y) # XLA adds the all-reduce
# PyTorch: wrap modules; hooks issue the collectives at run time.
import os, torch.distributed as dist
from torch.distributed.fsdp import fully_shard
dist.init_process_group("nccl")
torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
for block in model.blocks: # a transformer's block list
fully_shard(block)
fully_shard(model)In JAX you annotate how arrays are laid out across a device mesh and the XLA partitioner inserts the communication. When you want explicit control per device, jax.shard_map lets you write the per-shard program and call collectives such as jax.lax.psum yourself. In PyTorch you choose a wrapper, DistributedDataParallel or FSDP's fully_shard, and it registers hooks that issue NCCL calls during backward, overlapping them with computation; the FSDP article covers that lifecycle in detail. TensorFlow uses tf.distribute strategies such as MirroredStrategy for single-host data parallelism and MultiWorkerMirroredStrategy across hosts. All three end in the same NCCL operations; the NCCL collectives article explains their cost and topology.
Annotations make new layouts cheap to try, but a poor partitioner choice means debugging a compiled program. Wrappers are explicit, but combining tensor, pipeline and data parallelism means composing several of them correctly.
State, randomness and debugging
Randomness. PyTorch and TensorFlow keep global random state; JAX makes you split and pass keys. The JAX way is more work but makes dropout and data shuffling reproducible under jit and across devices by construction. In PyTorch, reproducibility requires seeding every process and accepting that some GPU kernels are non-deterministic unless you request deterministic algorithms, at a speed cost.
Debugging. Eager PyTorch is easiest. JAX offers jax.debug.print and a global switch to disable jit; TensorFlow has tf.config.run_functions_eagerly(True). In all three: debug eager, compile, then compare outputs numerically.
Buffer donation and memory. Because JAX returns new arrays, the old parameters would stay alive unless you tell the compiler it may reuse their memory with donate_argnums. Forgetting donation roughly doubles parameter and optimizer memory in a large model. PyTorch's in-place updates avoid this by default.
Worked example: porting a step and measuring it
Suppose a team trains a 350-million-parameter transformer in eager PyTorch on eight GPUs and profiling shows the GPU is idle between kernels about a third of the time. The pattern suggests CPU dispatch overhead and unfused element-wise work, not slow matrix multiplies. They have two realistic options.
- Compile in place. Wrap the model with
torch.compile, run with graph-break logging, fix the breaks in the hot path (usually data-dependent branches and non-tensor Python objects), and keep variable-length batches bucketed to a handful of shapes. This is a day of work and keeps every existing tool. - Rewrite in JAX. The whole step becomes one XLA program and sharding becomes annotations. This is weeks of work, requires a new checkpoint format and data pipeline, and pays off mainly if the team also wants TPU access or compiler-driven parallelism.
The right first move is almost always option one, measured properly: record step time for steps 20 to 120 after warm-up, check that the number of compiled graphs stops growing, and compare loss curves for the first thousand steps against the eager baseline to catch numerical drift. The GPU profiling guide shows how to confirm the idle gaps closed rather than moved.
Ecosystem and deployment
PyTorch dominates research code and open model releases, so most new architectures, attention kernels and inference servers appear there first. JAX is strongest where compiler-driven scaling matters and on TPUs, with Optax for optimizers and Orbax for checkpoints. TensorFlow remains widespread in existing production systems and in Keras-based teams; its on-device path is moving from tf.lite to the separate LiteRT project, so new mobile work should start there rather than on the deprecated module.
Failure modes
| Symptom | Framework | Cause | Fix |
|---|---|---|---|
| Every step slow, compile logs repeat | All | Changing input shapes force recompiles | Pad or bucket shapes; mark dynamic dimensions |
| torch.compile gives little speed-up | PyTorch | Graph breaks leave hot code eager | Log graph breaks; fullgraph=True on the hot module |
| ConcretizationTypeError | JAX | Python if/for on a traced value | Use lax.cond, lax.scan or jnp.where |
| Out of memory only under jit | JAX | Old buffers kept alive | donate_argnums for params and optimizer state |
| print or counter runs once | JAX, TF | Side effects happen at trace time | jax.debug.print, tf.print, return values |
| Retracing warnings, slow Keras fit | TF | Python scalars or new shapes as inputs | Pass tensors; set input_signature |
| Loss diverges after compiling | All | Different fusion or precision order | Compare against eager on a fixed batch; check dtypes |
| All-reduce never overlaps compute | PyTorch | Wrapping too coarse or too fine | Wrap per transformer block; profile NCCL streams |
Trade-offs and how to choose
| Need | Best fit | Why |
|---|---|---|
| Fast iteration on new research ideas | PyTorch | Eager by default, largest ecosystem |
| Whole-program compilation and sharding by annotation | JAX | jit plus sharding is the core design |
| TPU training | JAX | XLA is the native path |
| Existing TF or Keras production stack | TensorFlow or Keras 3 | Migration cost rarely pays back |
| One model definition, several engines | Keras 3 | Backend-agnostic layers |
| Custom GPU kernels in Python | PyTorch or JAX with Triton or Pallas | Both expose kernel-level escape hatches |
What to do next
- Write the three training steps above and run each on one GPU with the same batch; confirm the losses match to a few decimal places.
- Profile your current job and measure the fraction of time the GPU sits idle between kernels; that number tells you whether compilation will help.
- If you use PyTorch, compile the model, list graph breaks and fix those in the forward pass of your main block.
- Pad or bucket input shapes and verify the number of compiled graphs stays constant after warm-up.
- Try one sharding layout change: annotations in JAX, or per-block fully_shard in PyTorch, and measure peak memory and step time.
- Write down, for your team, which framework owns training and which owns serving, and keep that decision separate.