Kauldron tutorial: experiments as plain data with konfig, kontext, ktyping, and kd.train

A Coding Guide to Google Research’s Kauldron: Configs That Are Plain Data, Components Wired by String, and a JAX Trainer You Can Read End to End

If you’ve ever lost a morning to brittle imports, opaque overrides, or a shape error that hides the offending tensor, the tutorial will feel refreshingly direct: represent experiments as plain data, wire components by string paths, and check tensor shapes by name. The tutorial demonstrates four pieces that together aim to make experiment iteration faster and safer: konfig, kontext, ktyping, and kd.train (Trainer).

“Kauldron’s pitch is modularity: it is the glue, not the framework. Four pieces do the work:

konfig -> your experiment IS a Python call tree, and that tree is a plain dict

kontext -> parts are wired by string key paths, so they never import each other

ktyping -> Float[‘*b h w c’] checked at runtime, with named axes bound across args

kd.train -> Trainer: model + data + losses + metrics + optimizer, and nothing else”

How Kauldron’s pieces fit together

Think of Kauldron as a set of practical conventions plus a thin runner. The tutorial focuses on four modules:

  • konfig (from kauldron.konfig), configs are plain, nested data you can serialize; they describe Python calls and resolve into real objects when needed.
  • kontext (from kauldron.kontext), wiring is done with string key paths; components declare which keys they need and get values from a shared runtime context.
  • ktyping (kauldron.typing), runtime shape/type checks with named axes, using annotations like Float[‘*b n c’] and a decorator to validate binds across args/returns.
  • kd.train (Trainer), composes the resolved model, data pipeline, losses, metrics, and optimizer, and runs training with hooks for checkpointing and evaluation.

That combination is intentionally minimal. Libraries like Optax and Flax remain unchanged and are configured as data in konfig, then resolved at runtime. The tradeoff is fewer import-time dependencies and simpler overrides, but some issues, like mistyped keys or logical config errors, move to runtime.

konfig: experiments-as-data

Konfig turns an experiment into a tree of plain dicts (ConfigDicts). These can be serialized to JSON and later resolved into Python objects. The tutorial pins kauldron==1.4.2 for the notebook examples and shows typical patterns: expressing an optimizer as data (for example, optax.adam with a learning rate), lazy imports for notebook-defined classes (konfig.imports(lazy=True)), and reference chaining so downstream values follow upstream changes using cfg.ref.

Operational guardrail: konfig refuses to store resolved runtime objects inside a ConfigDict. If you try to stash a resolved Flax module into the config, konfig raises a ValueError immediately. That choice preserves serializability and predictable override semantics.

kontext: wiring by string key paths

Instead of importing upstream components directly, functions and modules declare the keys they expect. In the tutorial this looks like declaring an input:

  • kontext.Key = kontext.REQUIRED (used in losses/metrics/models)

At runtime, kontext.resolve_from_keyed_obj pulls values from a shared context using string paths such as “batch.image”, “preds.logits”, and nested forms like “preds.aux[0].pos”. The tutorial deliberately shows a mistyped path “batch.nope” to demonstrate the resulting KeyError and why runtime validation matters.

ktyping: named axes and clearer shape errors

Ktyping attaches semantic names to tensor axes and validates them at runtime. The tutorial demonstrates the pattern with a simple annotated function:

  • Float[‘*b n c’], Float[‘c d’] -> return Float[‘*b n d’]

Example shapes used in the notebook: features zeros with shape (2, 16, 8) and weights zeros with shape (8, 32) bind c=8 and d=32 to satisfy the annotations. When bindings fail, ktyping reports the named axis that conflicted, like c or d. That output is more actionable than an opaque tuple mismatch, especially in larger models with many tensor arguments.

Guardrails and tradeoffs

  • No resolved objects in configs: Konfig raises on attempts to store resolved runtime objects, preserving serializability but requiring you to keep runtime state separate from your ConfigDicts.
  • String wiring moves some checks to runtime: kontext reduces import coupling, but mistyped key paths become runtime KeyErrors. ktyping helps by catching shape issues early, and adding a small startup validation or smoke test reduces surprises.
  • Runtime checks have cost: named-axis checking is invaluable during development, but teams should consider when to disable or guard checks in hot production paths where JIT or pmapped performance matters.

The Trainer demo: readable end-to-end runs

The notebook assembles a small experiment and runs a complete training loop on synthetic, in-memory data (CPU only) to showcase ergonomics. Key elements from the tutorial:

  • Model: a small MLP that declares input via kontext key “batch.x”, with a default hidden size of 32 and a final Dense(1, name=”out”).
  • Data: a synthetic loader creating a train split of n=512 and eval split n=128; SEED_W is generated with np.random.default_rng(7).normal(size=(16, 1)).astype(“float32”); kd.data.InMemoryPipeline used with batch_size=32.
  • Optimizers and schedules: examples show optax.adam(learning_rate=0.003) as config data and an optax.chain with clip_by_global_norm(1.0), scale_by_adam(b2=0.99), scale_by_learning_rate(0.003). A warmup/cosine schedule example uses init_value=0.0, peak_value=1e-3, warmup_steps=100, and decay_steps set via cfg.ref.num_train_steps.
  • Trainer settings: example workdirs “/tmp/kauldron_tutorial” and “/tmp/kauldron_resume”; num_train_steps values shown include 300 and 200 in different demos.

Notebook-observed training behavior: in the CPU demo the loss fell from 1.51 to 0.005 over 300 steps. The author reports the run finished in roughly one second on a single CPU host. These are demo observations for this small, synthetic setup and should be treated as ergonomics evidence rather than production benchmarks.

Losses, metrics, and merging state

Custom losses and metrics follow expected shapes and state conventions. The tutorial shows:

  • A LogCosh loss subclassing kd.losses.Loss with preds and targets declared as kontext.Key REQUIRED.
  • A WithinTol metric subclassing kd.metrics.Metric, with an AutoState subclass using sum_field defaults (n_hit and n_total) so metric state merges by summing numerators and denominators across batches and devices, the statistically correct aggregation strategy when batch sizes vary.

Sweeps, checkpointing, and resume

Experiment ergonomics are highlighted by a five-variant sweep implemented as a simple for-loop that applies one-line config overrides. Overrides shown in the tutorial include:

  • cfg.model.hidden = 4 or = 128
  • cfg.optimizer.learning_rate = 0.1
  • cfg.optimizer = optax.sgd(0.05)

Each variant runs for 200 steps in the sweep example. Checkpoints are written using kd.ckpts.Checkpointer with save_interval_steps=100 and files named like checkpoints/ckpt_*. The resume demo trains to step 200 in “/tmp/kauldron_resume”, then rebuilds the Trainer requesting 300 steps. The second run resumes from step 200.

“Without the two-line replacement below, which uses jax’s own public dtype API and is a no-op on older jax, a Trainer raises AttributeError before it completes a single step.”, tutorial note about an etils/JAX compatibility workaround

The tutorial includes a two-line compatibility workaround addressing an etils/JAX mismatch observed with newer JAX and older etils. Because that patch is an operational detail with version-dependent implications, check the Kauldron repository and the notebook cell that contains the fix before applying anything locally; upstream fixes may already exist.

What the demo demonstrates, and what it doesn’t

Demonstrated value:

  • Experiments as plain, serializable ConfigDicts (konfig).
  • Loose coupling via string wiring (kontext) so parts never import each other.
  • Clearer shape diagnostics with named-axis runtime checks (ktyping).
  • An end-to-end Trainer that supports checkpointing, evaluation hooks, sweeps, and resume semantics.

Limits to keep in mind:

  • The notebook uses synthetic, in-memory data and CPU runs. It validates ergonomics and correctness patterns, not multi-host sharding, I/O throughput, or production TPU/GPU performance.
  • Compatibility with JAX, etils, flax, and optax versions may require attention. Follow the Kauldron repo’s compatibility notes.
  • String-based wiring reduces import coupling but requires disciplined tests and startup validation to avoid runtime KeyErrors in larger systems.

Where to go next

  • Kauldron repository and examples: github.com/google-research/kauldron/tree/main/examples
  • Documentation: kauldron.readthedocs.io
  • Try running the example CLI pattern shown in the tutorial: python -m kauldron.main –cfg=config.py –cfg.model.hidden=128

Practical checklist before you run the notebook

  • Pin kauldron==1.4.2 as used in the tutorial, and record the JAX and etils versions you install. Check the notebook cell that contains the compatibility note before applying any monkeypatches.
  • Run the synthetic CPU notebook first to confirm the workflow: printed steps, created workdir (for example /tmp/kauldron_tutorial), and checkpoint files.
  • Validate checkpoint/resume by training to a configured save step (e.g., 200), then restarting with a larger num_train_steps and confirming the run resumes at the saved step.
  • Before scaling, replace the synthetic loader with a single-host GPU run on a small real dataset to confirm basic sharding, JIT behavior, and step timing.

Key takeaways, questions you might be asking

  • How does Kauldron avoid importing components into each other?
    By wiring components with string key paths (kontext). Objects declare the keys they consume and the runner resolves those keys from a shared context at runtime.
  • Can I pass plain Optax and Flax objects without Kauldron-specific code?
    Yes. The tutorial demonstrates passing Optax and Flax objects as plain data in ConfigDicts and resolving them with konfig, neither Optax nor Flax needs Kauldron-aware modifications.
  • Will named-axis ktyping replace unit tests?
    No. ktyping reduces debugging time by making shape mismatches explicit but it is a runtime guard, not a substitute for behavioral tests and coverage across edge cases.
  • Does the notebook prove Kauldron is production-ready for TPU multi-host training?
    No. The notebook is an ergonomics demo (synthetic data, CPU). Validate multi-host sharding, data I/O, and checkpoint integrity on your target stack before migrating production jobs.
  • What should I watch for with JAX and etils versions?
    The tutorial notes a compatibility issue involving newer JAX and older etils that the notebook addresses with a two-line workaround. Confirm current compatibility in the Kauldron repo and prefer upstream fixes to local monkeypatches.
  • How much engineering effort to adopt Kauldron?
    A quick sanity check (running the tutorial notebook) can be done in a day. Converting a large existing codebase depends on how tightly components are import-coupled. Expect anything from a few days if components are modular to several weeks for deep integrations with custom checkpoint and serving systems.

Final note for engineering leaders

Kauldron is worth a short experiment if your team spends time firefighting imports, opaque overrides, or cryptic shape errors. Its core idea, make experiments a tree of plain data, wire by key paths, and validate axes by name, targets developer velocity and clearer failure modes more than raw training throughput.

Staged adoption plan (practical):

  1. Run the tutorial notebook end-to-end on CPU and confirm printed steps, loss behavior, checkpoints, and resume semantics.
  2. Swap the synthetic loader for a small real dataset on a single GPU host; validate step time, JIT behavior, and checkpoint integrity.
  3. Validate distributed behavior on a low-cost multi-host testbed (small batch sizes) to exercise sharding, metric merges, and checkpointing across hosts.

If those stages pass, you gain a reproducible, override-friendly experiment surface that often shortens iteration cycles. If they reveal friction, the gaps are obvious: wiring or config mismatches, version incompatibilities, or integration points with existing infra (schedulers, logging, model export) that need engineering work. That makes the risk easy to quantify and plan for.

For authoritative details and the exact notebook cells referenced (including the compatibility workaround), consult the Kauldron repo and readthedocs pages listed above before applying any local fixes or rolling Kauldron into production workflows.