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.

Advertisement

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 loss

TensorFlow 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.

Advertisement

Down to the GPU: dispatch, compilers and kernels

From Python to GPU kernels: three execution modelsPyTorchPython codeeager ops run as calledCapturetorch.compile: Dynamo, FXCompilerInductor (optional)KernelscuBLAS, cuDNN, TritonCollectivesNCCL via torch.distributedJAXPython codepure functions + pytreesCapturejit tracing to jaxprCompilerXLA via StableHLOKernelsXLA fused kernels, cuBLASCollectivesXLA-inserted NCCLTensorFlow / KerasPython codeKeras model / tf.functionCapturetracing to a TF graphCompilerXLA when jit_compile=TrueKernelsTF op kernels or XLACollectivesNCCL via tf.distribute
Each framework captures Python differently, but all end at the same place: vendor libraries such as cuBLAS and cuDNN for large matrix work, generated kernels for the glue between them, and NCCL for communication between GPUs.

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.

  1. 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.
  2. 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

SymptomFrameworkCauseFix
Every step slow, compile logs repeatAllChanging input shapes force recompilesPad or bucket shapes; mark dynamic dimensions
torch.compile gives little speed-upPyTorchGraph breaks leave hot code eagerLog graph breaks; fullgraph=True on the hot module
ConcretizationTypeErrorJAXPython if/for on a traced valueUse lax.cond, lax.scan or jnp.where
Out of memory only under jitJAXOld buffers kept alivedonate_argnums for params and optimizer state
print or counter runs onceJAX, TFSide effects happen at trace timejax.debug.print, tf.print, return values
Retracing warnings, slow Keras fitTFPython scalars or new shapes as inputsPass tensors; set input_signature
Loss diverges after compilingAllDifferent fusion or precision orderCompare against eager on a fixed batch; check dtypes
All-reduce never overlaps computePyTorchWrapping too coarse or too fineWrap per transformer block; profile NCCL streams

Trade-offs and how to choose

NeedBest fitWhy
Fast iteration on new research ideasPyTorchEager by default, largest ecosystem
Whole-program compilation and sharding by annotationJAXjit plus sharding is the core design
TPU trainingJAXXLA is the native path
Existing TF or Keras production stackTensorFlow or Keras 3Migration cost rarely pays back
One model definition, several enginesKeras 3Backend-agnostic layers
Custom GPU kernels in PythonPyTorch or JAX with Triton or PallasBoth expose kernel-level escape hatches

What to do next

  1. 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.
  2. Profile your current job and measure the fraction of time the GPU sits idle between kernels; that number tells you whether compilation will help.
  3. If you use PyTorch, compile the model, list graph breaks and fix those in the forward pass of your main block.
  4. Pad or bucket input shapes and verify the number of compiled graphs stays constant after warm-up.
  5. Try one sharding layout change: annotations in JAX, or per-block fully_shard in PyTorch, and measure peak memory and step time.
  6. Write down, for your team, which framework owns training and which owns serving, and keep that decision separate.
Key takeaway: PyTorch, JAX and TensorFlow differ mainly in execution model: PyTorch runs eagerly and compiles optionally, JAX traces pure functions into whole XLA programs, and TensorFlow traces Python into graphs that XLA can compile. All three end in the same vendor kernels and NCCL collectives; the differences are in fusion, overhead, how shapes and control flow behave under tracing, and who inserts communication. Pick the compilation boundary and ecosystem your team needs, then measure.