In this tutorial, we implement Kauldron, the JAX training library from Google Research that describes itself as optimized for research velocity and modularity, and we take those two words literally by testing what they actually buy us. We install it, then spend the first half of the notebook on the three mechanisms that make Kauldron different from a stack of Flax and Optax: konfig, which turns an experiment into a tree of plain dictionaries that round-trip through JSON; kontext, which wires components together with string key paths so a loss never imports the model it scores; and the runtime shape checker, whose named axes bind across arguments and report what they were bound to when something does not match. We then write a custom loss and a custom metric in the shape the framework expects, train a real Trainer on synthetic in-memory data with no download and no accelerator, and monitor an inner layer of the model without editing the model. We finish by running a five-variant sweep in which every experiment differs by one config line, and by letting a training run checkpoint itself and resume where it stopped.
import os
import sys
import json
import textwrap
import traceback
import subprocess
RESULTS = {}
def banner(title):
print("\n" + "=" * 78)
print(title)
print("=" * 78)
def section(name):
def wrap(fn):
def run(*a, **kw):
banner(name)
try:
out = fn(*a, **kw)
RESULTS[name] = out if isinstance(out, str) else "ok"
return out
except Exception as e:
RESULTS[name] = f"SKIPPED / FAILED -> {type(e).__name__}: {e}"
print(f"\n[!] {name} did not complete: {type(e).__name__}: {e}")
traceback.print_exc(limit=3)
return None
return run
return wrap
banner("0. Install Kauldron, and the one compatibility patch you need today")
subprocess.run([sys.executable, "-m", "pip", "install", "-q", "kauldron==1.4.2"], check=True)
import jax
from etils.enp import array_spec as _array_spec
# jax >= 0.10.1 moved `jax._src.prng`, but etils <= 1.14.0 still reaches for it whenever it
# inspects an array's dtype. Kauldron calls that code on every batch, so without this two-line
# patch a Trainer raises AttributeError before it finishes a single step. The replacement uses
# jax's own public dtype API and is a no-op on older jax.
if not hasattr(jax._src, "prng"):
_array_spec._is_jax_random_dtype = lambda dt: jax.dtypes.issubdtype(dt, jax.dtypes.prng_key)
import numpy as np
import optax
import flax
from flax import linen as nn
import kauldron
from kauldron import kd, konfig, kontext
from kauldron.typing import Float, typechecked
print(f" kauldron {kauldron.__version__} | jax {jax.__version__} | flax {flax.__version__}"
f" | optax {optax.__version__}")
print(f" devices: {jax.devices()}")
print("\n Kauldron's pitch is modularity: it is the glue, not the framework. Four pieces do the work:")
print(" konfig -> your experiment IS a Python call tree, and that tree is a plain dict")
print(" kontext -> parts are wired by string key paths, so they never import each other")
print(" ktyping -> Float['*b h w c'] checked at runtime, with named axes bound across args")
print(" kd.train -> Trainer: model + data + losses + metrics + optimizer, and nothing else")
print("\n Everything below runs on a CPU runtime with no dataset download: the data is synthetic.")
We install Kauldron and apply the one compatibility patch the current release combination needs. jax 0.10.1 moved the private module jax._src.prng, and etils up to 1.14.0 still reaches for it whenever it inspects an array’s dtype, which is a code path Kauldron runs on every batch. Without the two-line replacement below, which uses jax’s own public dtype API and is a no-op on older versions, a Trainer raises AttributeError before it completes a single step. With it in place we import the four pieces that do the work: konfig for the config system, kontext for the wiring, the typing module for runtime shape checks, and kd.train for the Trainer itself. Everything afterwards runs on a CPU runtime, because the only dataset in this notebook is one we generate.
@section("1. A config is a call tree, and a call tree is a dict")
def config_is_a_dict():
with konfig.imports():
import optax as coptax # looks like optax, builds ConfigDict instead
cfg = coptax.adam(learning_rate=0.003)
print(f" cfg = {cfg}")
print(f" type = {type(cfg).__name__}")
print(f" __qualname__ = {cfg.__qualname__!r} <- the call, stored as data")
cfg.learning_rate = 1e-4 # configs are mutable
optimizer = konfig.resolve(cfg) # ...until you resolve them
print(f" after cfg.learning_rate = 1e-4 -> resolve() gives {type(optimizer).__name__}")
print("\n An arbitrarily complex optimizer is still just nested dicts:")
chain = coptax.chain(
coptax.clip_by_global_norm(1.0),
coptax.scale_by_adam(b2=0.99),
coptax.scale_by_learning_rate(0.003),
)
as_json = json.dumps(json.loads(chain.to_json()), indent=2)
print(textwrap.indent(as_json, " "))
rebuilt = konfig.resolve(konfig.ConfigDict(json.loads(chain.to_json())))
print(f" JSON -> ConfigDict -> resolve() -> {type(rebuilt).__name__}")
print(" optax has no idea konfig exists. No base class, no registry, no decorator.")
return f"optax.chain -> JSON -> {type(rebuilt).__name__}"
config_is_a_dict()
We start with konfig, because it is the piece the rest of the library is built on. Inside a konfig.imports() block, importing optax gives us something that looks and autocompletes like optax but builds configuration instead of objects, so optax.adam(learning_rate=0.003) returns a ConfigDict holding the qualified name of the call and its arguments rather than an optimizer. That config is mutable until konfig.resolve turns it into the real thing, and because it is only nested dictionaries, an arbitrarily complex optax.chain serialises to JSON and comes back as a working optimizer. The important part is what optax had to do to support this: nothing. There is no base class, no registry, and no decorator anywhere in optax, and the same applies to any library we configure this way.
@section("2. cfg.ref: change one number, everything downstream follows")
def config_references():
with konfig.imports():
import optax as coptax
from kauldron import kd as ckd
cfg = ckd.train.Trainer()
cfg.num_train_steps = 1000
cfg.schedules = {
"lr": coptax.warmup_cosine_decay_schedule(
init_value=0.0, peak_value=1e-3, warmup_steps=100,
decay_steps=cfg.ref.num_train_steps, # <- a reference, not the value 1000
)
}
at_1000 = konfig.resolve(cfg.schedules["lr"])
cfg.num_train_steps = 200 # one edit...
at_200 = konfig.resolve(cfg.schedules["lr"]) # ...and the schedule already knows
print(f" {'progress':>10s} {'lr @ 1000 steps':>18s} {'lr @ 200 steps':>16s}")
for frac in (0.1, 0.5, 0.9):
print(f" {frac:>9.0%} {float(at_1000(int(1000*frac))):>18.6f}"
f" {float(at_200(int(200*frac))):>16.6f}")
print("\n Without .ref the schedule would have frozen 1000 into itself, and a sweep over")
print(" num_train_steps would have silently trained on the wrong decay curve.")
return (f"lr at 90% of training: {float(at_1000(900)):.6f} (1000 steps)"
f" vs {float(at_200(180)):.6f} (200 steps)")
config_references()
Configuration systems usually go wrong when one value is needed in several places, and Kauldron’s answer is cfg.ref. We point a warmup-cosine schedule’s decay_steps at cfg.ref.num_train_steps rather than 1000, then change num_train_steps to 200 and resolve the schedule again. The learning rate curve reshapes itself, because the config stored a reference rather than a copy of the value. Without that indirection the schedule would have frozen 1000 into itself, and a sweep over the number of training steps would have quietly trained every variant on the wrong decay curve, which is the kind of bug that produces a plausible number and no error.
@section("3. kontext: parts are wired by string, so they never import each other")
def kontext_keys():
import dataclasses
ctx = {
"batch": {"image": np.zeros((4, 8, 8, 3)), "label": np.arange(4)},
"preds": {"logits": np.ones((4, 10)), "aux": [{"pos": np.zeros(3)}]},
}
print(" a context is just nested data; a key path reaches into it:")
for path in ["batch.image", "preds.logits", "preds.aux[0].pos"]:
print(f" {path:22s} -> {kontext.get_by_path(ctx, path).shape}")
try:
kontext.get_by_path(ctx, "batch.nope")
except KeyError as e:
print(f" {'batch.nope':22s} -> KeyError: {str(e)[:96]}...")
@dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
class MeanGap:
preds: kontext.Key = kontext.REQUIRED # these ARE the wiring
targets: kontext.Key = kontext.REQUIRED
def __call__(self, *, preds, targets):
return float(abs(np.asarray(preds).mean() - np.asarray(targets).mean()))
metric = MeanGap(preds="preds.logits", targets="batch.label")
kwargs = kontext.resolve_from_keyed_obj(ctx, metric)
print(f"\n MeanGap declared preds={metric.preds!r}, targets={metric.targets!r}")
print(f" resolved to kwargs: {{{', '.join(f'{k}: {v.shape}' for k, v in kwargs.items())}}}")
print(f" value = {metric(**kwargs)}")
print("\n MeanGap never imported the model and the model never heard of MeanGap. Point the")
print(" same metric at 'preds.aux[0].pos' and nothing but that string changes.")
return f"MeanGap(preds='preds.logits', targets='batch.label') = {metric(**kwargs)}"
kontext_keys()
kontext is how Kauldron connects components that know nothing about each other. A context is ordinary nested data, and a key path such as batch.image or preds.aux[0].pos reaches into it, resolving dictionary keys, attributes and list indices alike, and raising a KeyError that lists what was actually available when it cannot. Any object can declare its inputs by annotating fields as kontext.Key, and resolve_from_keyed_obj then pulls exactly those paths out of the context and hands them over as keyword arguments. We build a small metric this way and point it at a model’s outputs: the metric never imports the model, the model never hears of the metric, and redirecting the metric at a different tensor is a change to one string.
@section("4. ktyping: named axes, checked at runtime, bound across arguments")
def shape_checking():
@typechecked
def project(features: Float["*b n c"], weights: Float["c d"]) -> Float["*b n d"]:
return jax.numpy.einsum("...c,cd->...d", features, weights)
out = project(jax.numpy.zeros((2, 16, 8)), jax.numpy.zeros((8, 32)))
print(f" project(f32[2 16 8], f32[8 32]) -> {out.shape} c bound to 8, d bound to 32")
print("\n now break it: c is bound to 8 by the first argument, so 5 cannot also be c")
try:
project(jax.numpy.zeros((2, 16, 8)), jax.numpy.zeros((5, 32)))
except Exception as e:
print(textwrap.indent(str(e), " "))
print("\n 'Inferred Dims' is the part worth having: it reports what each axis name was already")
print(" bound to, so a mismatch names the axis instead of printing two anonymous shapes.")
return "mismatch named the axis: c already bound to 8, got 5"
shape_checking()
Kauldron’s typing module checks array shapes at runtime using named axes. We annotate a function with Float[‘*b n c’] and Float[‘c d’], and the decorator binds each axis name the first time it sees it, then enforces that binding everywhere else in the signature, including the return value. When we deliberately pass an incompatible second argument, the error does the thing that matters: alongside the actual shapes it prints an Inferred Dims block showing that c had already been bound to 8, so the failure names the axis that disagreed instead of leaving us to compare two anonymous tuples. On a model with several tensors in flight this is the difference between a one-line fix and a debugging session.
import dataclasses
@dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
class LogCosh(kd.losses.Loss):
"""log(cosh(err)): quadratic near zero, linear in the tails. ~30 lines less than raw Flax."""
preds: kontext.Key = kontext.REQUIRED
targets: kontext.Key = kontext.REQUIRED
@typechecked
def get_values(self, preds: Float["*a"], targets: Float["*a"]) -> Float["*a"]:
return jax.numpy.log(jax.numpy.cosh(preds - targets))
@dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
class WithinTol(kd.metrics.Metric):
"""Fraction of predictions landing within `tol` of the target, over every batch seen."""
preds: kontext.Key = kontext.REQUIRED
targets: kontext.Key = kontext.REQUIRED
tol: float = 0.25
@flax.struct.dataclass
class State(kd.metrics.AutoState):
# sum_field() marks a value that is ADDED when two states merge. Keeping the numerator
# and the denominator apart is what makes the pooled result exact.
n_hit: Float[""] = kd.metrics.sum_field(default=0.0)
n_total: Float[""] = kd.metrics.sum_field(default=0.0)
def compute(self) -> Float[""]:
# Return a jax scalar, like the built-in states do: the metric writer that
# `trainer.train()` logs through does not accept a bare numpy scalar.
total = jax.numpy.maximum(jax.numpy.asarray(self.n_total), 1.0)
return jax.numpy.asarray(self.n_hit) / total
@typechecked
def get_state(self, preds: Float["*a"], targets: Float["*a"]) -> "WithinTol.State":
hit = (jax.numpy.abs(preds - targets) < self.tol).astype("float32")
return self.State(n_hit=hit.sum(), n_total=jax.numpy.asarray(hit.size, "float32"))
@section("5. A custom loss and a custom metric, in the shape Kauldron expects")
def custom_loss_and_metric():
rng = np.random.default_rng(0)
p = jax.numpy.asarray(rng.normal(size=(8, 4)).astype("float32"))
t = jax.numpy.asarray(rng.normal(size=(8, 4)).astype("float32"))
loss = LogCosh(preds="preds.y", targets="batch.y")
print(f" {'LogCosh(preds, targets)':34s} {float(loss(preds=p, targets=t)):.6f}")
print(f" {'same loss, weight=0.5':34s} "
f"{float(LogCosh(preds='a', targets='b', weight=0.5)(preds=p, targets=t)):.6f} <- exactly half")
print(f" {'builtin kd.losses.L2':34s} {float(kd.losses.L2(preds='a', targets='b')(preds=p, targets=t)):.6f}")
print("\n A metric is not a number, it is a State that merges. Watch why that matters when")
print(" the last batch of an epoch is smaller than the rest:")
metric = WithinTol(preds="preds.y", targets="batch.y", tol=0.5)
big = metric.get_state(preds=p[:6], targets=t[:6])
small = metric.get_state(preds=p[6:], targets=t[6:])
merged = big.merge(small)
fo