Pytrees
A pytree is any nest of dicts, lists and tuples with arrays at the ends — and JAX can reach every array in one sweep, however deep the nesting.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A pytree is any nested arrangement of dicts, lists and tuples whose end points are arrays — and JAX knows how to visit every end point in one instruction.
Think of a hostel building. Floors contain rooms, rooms contain cupboards, cupboards contain shelves. The warden says one sentence — "dust every shelf" — and the cleaning staff visit all of them, whatever floor, whatever room. Nobody lists the shelves one by one.
The building is the pytree: the nesting of floors and rooms is the structure, and the shelves — the actual arrays — are called leaves. JAX's tree tools are the cleaning staff.
Why it exists
A real model's parameters are not one array. They are dozens: each layer has its weights and biases, layers group into blocks, blocks into the model. The natural home for that is nested dicts — readable, named, organised.
But training needs to do things to all parameters at once: nudge each by its gradient, zero them, measure their size. Writing a loop through every layer, every time, for every operation, would bury the actual idea. The pytree machinery lets one line mean "do this to every leaf". Better still: "combine these two matching trees, leaf by leaf", which is exactly what a gradient update is.
How it works
params: grads: one line:
{ layer1: {w: ▓, b: ▓}, { layer1: {w: ░, b: ░}, new = map(p - 0.1*g
layer2: {w: ▓, b: ▓} } layer2: {w: ░, b: ░} } over both trees)
same structure ────────────── same structure ────▶ same structure outThe key fact: grad returns gradients in the same tree shape as the parameters. Matching trees combine leaf-for-leaf, like two stacks of forms filled in the same order.
A real example you have seen
Folders on your computer. "Find all photos in this folder and its subfolders" works no matter how messy your nesting is. The search visits every file without you opening each folder. Tree operations are that search, pointed at arrays instead of photos.
Remember this
- Pytree = nested dicts/lists/tuples; the arrays at the ends are leaves.
jax.tree.mapapplies a function to every leaf, or combines matching trees.- Gradients arrive in the same tree shape as parameters — that is what makes updates one line.
What to learn next
- Random numbers and PRNG keys — the last plumbing before building real models.
- Building models with Flax — modules whose parameters are exactly these trees.
- Optax — tree-shaped optimizer states in action.
Developer — Code and libraries.
Setup
pip install jaxOutputs verified with jax 0.6.2, CPU. (jax.tree.map is the current spelling; older code says jax.tree_util.tree_map — same function.)
Parameters as a tree, updated in one line
import jax
params = {
"layer1": {"w": jax.numpy.ones((2, 3)), "b": jax.numpy.zeros(3)},
"layer2": {"w": jax.numpy.ones((3, 1)), "b": jax.numpy.zeros(1)},
}
# One call reaches every array in the nest, however deep
shapes = jax.tree.map(lambda leaf: leaf.shape, params)
print(shapes)
leaves = jax.tree.leaves(params)
print("number of leaves:", len(leaves))
print("total parameters:", sum(leaf.size for leaf in leaves))
# A gradient update touches every leaf without naming any of them
grads = jax.tree.map(lambda leaf: leaf * 0 + 0.5, params) # pretend gradients
updated = jax.tree.map(lambda p, g: p - 0.1 * g, params, grads)
print("layer1 w after update:", updated["layer1"]["w"][0]){'layer1': {'b': (3,), 'w': (2, 3)}, 'layer2': {'b': (1,), 'w': (3, 1)}}
number of leaves: 4
total parameters: 13
layer1 w after update: [0.95 0.95 0.95]The walkthrough
jax.tree.map(f, tree) applies f to each leaf and rebuilds the same structure around the results. The shapes printout keeps the dict nesting — instantly readable as a model summary. This shape-audit one-liner is worth memorising; it is the first thing to print when any JAX code misbehaves.
jax.tree.leaves flattens to a plain list, structure discarded — right for reductions like counting, wrong for anything that must go back into the model. Its inverse pair is jax.tree.flatten / jax.tree.unflatten, which keep the structure ticket so you can rebuild.
The two-tree form is the training step. jax.tree.map(lambda p, g: p - 0.1 * g, params, grads) walks both trees in lockstep — leaf pairs meet, combine, and a new tree returns. Every JAX training loop has this line, or hides it inside Optax. Check the arithmetic in the output: 1.0 − 0.1 × 0.5 = 0.95.
Everything accepts pytrees, not only tree.map. grad differentiates functions of pytrees, jit compiles them, vmap batches them — arguments and returns can be arbitrary nests everywhere in JAX. This is why the library needs no Parameter class and no Module base: plain containers are the parameter format. Structures you register yourself (dataclasses) can join too.
Common mistakes
Mismatched structures in multi-tree map. If grads has a leaf params lacks — or a list where a tuple should be — you get ValueError: ... tree structure mismatch (wording varies). Diagnose with jax.tree.structure(a) == jax.tree.structure(b). The subtle cousin: None values, which are treated as empty structure, not a leaf, and silently vanish from maps.
Forgetting a tuple is a tree too. Return (loss, params) from a function and tree operations will happily descend into it. If you meant "treat this whole thing as one value", the is_leaf= argument draws the boundary.
Loops over .items() doing what map does. Hand-walking nested dicts works until the nesting changes depth, then breaks. tree.map is depth-proof — code written for layer1/layer2 survives becoming block/layer/sublayer untouched.
Assuming an ordering. Dict leaves flatten in sorted-key order, not insertion order. Code that pairs tree.leaves(params) with tree.leaves(grads) by position works only because both sort the same way — rely on matched structure, not on remembered order.
Try it yourself
Compute the largest absolute value across all parameters in updated — flatten with jax.tree.leaves, then one max. Then write clip = lambda t: jax.tree.map(lambda x: jax.numpy.clip(x, -0.9, 0.9), t) and confirm shapes survive: jax.tree.map(lambda a: a.shape, clip(updated)).
What to learn next
- Random numbers and PRNG keys — the last plumbing before building real models.
- Building models with Flax — modules whose parameters are exactly these trees.
- Optax — tree-shaped optimizer states in action.
Researcher — Mathematics and papers.
The formal object
A pytree is defined inductively: a leaf (any object not registered as a container), or a registered container node holding pytrees as children. Registered by default: list, tuple, dict (keys sorted for canonical order), namedtuple, None (an empty node). flatten computes the pair $(\ell, s)$ — the leaf list $\ell$ in deterministic traversal order and the static structure $s$ (a PyTreeDef); unflatten(s, \ell)$ inverts it. tree.map` is flatten → zip → apply → unflatten, $O(|V|)$ in tree nodes per call.
The structure/leaf split is precisely how transformations handle containers: jit embeds $s$ in the compilation cache key and passes only $\ell$ (as typed arrays) into the traced program. Consequences worth internalising: changing container structure between calls retraces even when the arrays match (the cache-key churn of jit); structure is compile-time constant, so Python code inspecting it inside a jitted function is free; and custom containers registered via register_pytree_node (or flax.struct.dataclass) must place static configuration in the structure and dynamic arrays in the leaves, or every config change becomes a retrace.
Why containers-as-calling-convention matters
Frameworks with stateful modules (PyTorch nn.Module, Keras layers) thread parameters implicitly — objects own arrays, and the framework must provide a parallel API to traverse them (.parameters(), state_dict). JAX externalises the traversal as a first-class data-structure algebra, which is what makes the functional style workable at scale: optimizer states, gradients, EMA copies, and checkpoints are all trees of the same shape, manipulated by one operator. Libraries stack on this: Flax modules emit pytrees of params; Optax optimizer states are pytrees mirroring them; checkpointing (Orbax) serialises trees by path.
The path utilities (jax.tree_util.tree_map_with_path) expose per-leaf key paths — the mechanism behind readable parameter names in checkpoints and per-layer learning-rate maps (optax.multi_transform selects by path predicate).
The concept has no single origin paper; it formalises the "nested structure" conventions of Lisp-era mapcar lineage. The JAX documentation's pytree specification is the authoritative reference; the design rationale appears in Frostig, Johnson, Leary (2018), Compiling machine learning programs via high-level tracing.
References
- JAX documentation, Pytrees — the specification, including custom node registration.
- Frostig, Johnson, Leary (2018), Compiling machine learning programs via high-level tracing.
- Heek et al. (2020–), Flax: a neural network library and ecosystem for JAX — pytrees as the parameter format in practice.
What to learn next
- Random numbers and PRNG keys — the last plumbing before building real models.
- Building models with Flax — modules whose parameters are exactly these trees.
- Optax — tree-shaped optimizer states in action.